diff --git a/Cargo.lock b/Cargo.lock index 4ce55f3b..31b26d02 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -9,6 +9,7 @@ dependencies = [ "agent-client-protocol-derive", "agent-client-protocol-schema", "agent-client-protocol-test", + "async-channel", "async-io", "async-process", "blocking", @@ -110,11 +111,11 @@ dependencies = [ "agent-client-protocol", "async-stream", "axum", + "base64 0.23.1", "futures", - "futures-concurrency", - "rustc-hash", + "hmac", "serde_json", - "thiserror", + "sha2", "tokio", "tracing", "uuid", @@ -140,8 +141,7 @@ dependencies = [ [[package]] name = "agent-client-protocol-schema" version = "1.9.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1d595a79e1665d02c91c7dd1bedd118501e75cc0a3d5b1651b5150ec4c5c1e0d" +source = "git+https://github.com/agentclientprotocol/agent-client-protocol?rev=e5c36d2671fd355f983533bc83b5feb7981d25a6#e5c36d2671fd355f983533bc83b5feb7981d25a6" dependencies = [ "anyhow", "derive_more", @@ -878,6 +878,7 @@ checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" dependencies = [ "block-buffer 0.10.4", "crypto-common 0.1.7", + "subtle", ] [[package]] @@ -1203,6 +1204,15 @@ version = "0.4.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7f24254aa9a54b5c858eaee2f5bccdb46aaf0e486a595ed5fd8f86ba55232a70" +[[package]] +name = "hmac" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" +dependencies = [ + "digest 0.10.7", +] + [[package]] name = "http" version = "1.5.0" @@ -2536,6 +2546,17 @@ dependencies = [ "digest 0.11.3", ] +[[package]] +name = "sha2" +version = "0.10.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest 0.10.7", +] + [[package]] name = "sharded-slab" version = "0.1.7" diff --git a/Cargo.toml b/Cargo.toml index e62bf28d..db895a06 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -35,7 +35,8 @@ agent-client-protocol-trace-viewer = { path = "src/agent-client-protocol-trace-v yopo = { package = "agent-client-protocol-yopo", path = "src/yopo" } # Protocol -agent-client-protocol-schema = { version = "=1.9.1", default-features = false, features = ["tracing"] } +# Draft cross-repository validation; replace with the released schema before publishing. +agent-client-protocol-schema = { git = "https://github.com/agentclientprotocol/agent-client-protocol", rev = "e5c36d2671fd355f983533bc83b5feb7981d25a6", default-features = false, features = ["tracing"] } # Core async runtime tokio = { version = "1.52", default-features = false } @@ -44,6 +45,7 @@ tokio-util = { version = "0.7", features = ["compat"] } async-tungstenite = { version = "0.35.0", default-features = false, features = ["tokio-rustls-webpki-roots"] } # Serialization +base64 = "0.23" serde = { version = "1.0", features = ["derive", "rc"] } serde_json = { version = "1", features = ["preserve_order", "raw_value"] } schemars = { version = "1.0", features = ["derive"] } @@ -71,6 +73,7 @@ url = "2.5" async-io = "2" async-process = "2" async-stream = "0.3.6" +async-channel = "2" blocking = "1" chrono = "0.4" futures = "0.3.32" diff --git a/README.md b/README.md index f62a2572..6dafef29 100644 --- a/README.md +++ b/README.md @@ -34,6 +34,11 @@ attaches one while forking. Successful v2 attachments remain active for the connection lifetime, and all three builders expose `on_proxy_session_start` to forward proxied setup without coupling later session events to that response. +The native transport targets MCP 2026-07-28: requests carry their own metadata +and logical IDs, with request-scoped notifications and cancellation rather than +an MCP connection lifecycle. See [Native MCP-over-ACP](./md/mcp-over-acp.md) +for the direct rmcp example and current resource-limit caveats. + **Proxy orchestration** - [`agent-client-protocol-conductor`](./src/agent-client-protocol-conductor/) – Binary and library that manages chains of proxy components. diff --git a/justfile b/justfile index 6dd8c25d..314805e7 100644 --- a/justfile +++ b/justfile @@ -1,9 +1,12 @@ +# Keep file-based snapshots inside this checkout, even in nested worktrees. +export CARGO_WORKSPACE_DIR := justfile_directory() + # Build binaries needed for integration tests prep-tests: cargo build -p agent-client-protocol-conductor --all-features cargo build -p agent-client-protocol-test --bin testy --all-features cargo build -p agent-client-protocol-test --bin mcp-echo-server --example arrow_proxy --all-features -# Run all tests (requires prep-tests first) -test: prep-tests - cargo test --all --workspace --all-features +# Run all tests, or pass a test-name filter / cargo test arguments. +test *args: prep-tests + cargo test --all --workspace --all-features {{args}} diff --git a/md/SUMMARY.md b/md/SUMMARY.md index c51ce8a0..31ed83c7 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) @@ -31,6 +32,7 @@ # Reference +- [Migrating the Native MCP Transport](./migration-stateless-mcp.md) - [Migrating the rmcp Integration to v4](./migration-rmcp-v4.md) - [Migrating to v2.0](./migration_v2.0.md) - [Migrating to v0.11](./migration_v0.11.x.md) diff --git a/md/http-transport.md b/md/http-transport.md index c8584e76..05316355 100644 --- a/md/http-transport.md +++ b/md/http-transport.md @@ -59,10 +59,36 @@ active session, clients should also open: - `Acp-Connection-Id: ` - `Acp-Session-Id: ` -Open a session stream before sending methods such as `session/prompt`, -`session/load`, `session/resume`, or other session-scoped requests. When a -`session/new` or `session/fork` response returns a new `sessionId`, open an SSE -stream for that returned session before expecting updates or responses for it. +For a session not yet used on this HTTP connection, send its first +session-scoped POST (for example, `session/load` or `session/resume`) and wait +for `202 Accepted` before opening its GET. The POST registers the bounded +session mailbox before it is admitted to the agent. History and other output +can then queue until the GET attaches; an unknown-session GET returns 409 and +does not allocate a mailbox. A batch registers all its session mailboxes before +returning 202. + +The server reserves inbound queue capacity before publishing that metadata. +A rejected POST therefore cannot roll back a mailbox another accepted POST +has adopted. Registration failure or cancellation before publication releases +the reserved queue slot without leaving partial metadata. + +Keep the connection stream and existing session streams running during this +setup so callback responses and cancellation can still progress. `HttpClient` +does this automatically. When a `session/new` or `session/fork` response returns +a new `sessionId`, its mailbox is already registered and the client can open +the corresponding GET immediately. + +Frame and pending-work limits apply to HTTP and WebSocket traffic. Session +metadata retains a measured charge for its ID, not the entire opening request. +Consuming an SSE/WebSocket frame releases its payload charge before waiting +for the next frame. A WebSocket frame rejected by admission terminates that +connection rather than silently losing input or waiting forever to drain a +still-live agent. + +A fatal outbound routing failure, such as mailbox overflow, closes the whole +connection and reclaims its agent and metadata even if the agent emits no more +messages. This is distinct from an ordinary SSE disconnect, which can reconnect +to the existing mailbox. ## Features @@ -88,7 +114,13 @@ agent-client-protocol-http = { version = "...", features = ["client", "server"] `$/cancel_request` is connection-scoped. The HTTP transport does not apply `Acp-Session-Id` to cancellation notifications, and routes outgoing cancellation notifications over the connection stream rather than a session -stream. +stream. Cancellation is advisory: a `202 Accepted` for the cancellation POST +does not complete the original request. Pending response routing and its +bounded metadata charge stay in place until a terminal response, POST failure, +or connection teardown. The original request can still succeed after +cancellation, including a `session/new` or `session/fork` that opens a new +session stream. If the peer never sends a response, that request continues to +occupy pending-request capacity until the connection closes. WebSocket connections can carry cancellation at any point after the socket is open. With HTTP + SSE, cancellation can be sent after `initialize` completes and diff --git a/md/mcp-bridge.md b/md/mcp-bridge.md index 7caddd40..af412d3c 100644 --- a/md/mcp-bridge.md +++ b/md/mcp-bridge.md @@ -1,14 +1,18 @@ -# MCP-over-ACP Compatibility Bridge +# Stateless MCP-over-ACP HTTP Adapter `agent-client-protocol-polyfill::mcp_over_acp::McpOverAcpPolyfill` adapts the -native ACP MCP transport for a final agent that accepts HTTP MCP -servers. MCP adaptation is explicit and is not built into the conductor. +native ACP MCP transport for a final agent with an MCP 2026-07-28 HTTP client. +MCP adaptation is explicit and is not built into the conductor. There is no +fallback to older MCP revisions. The component-facing side of the bridge always uses the opt-in native protocol: - Servers are declared as `McpServer::Acp` with a `serverId`. -- Connections use `mcp/connect`, `mcp/message`, and `mcp/disconnect`. -- `mcp/disconnect` is a request with a response. +- Each operation uses `mcp/message` with `serverId` and a logical `requestId`. +- The provider sends notifications for that operation; the final ACP response + carries its MCP result or error. +- ACP request cancellation stops only that operation. There is no MCP + initialize/connect/disconnect or session-header exchange. The SDK-local underscore-prefixed method family and HTTP declarations with a special URL scheme have been retired. The polyfill now translates native @@ -88,17 +92,20 @@ support and rejects any native declaration that is nevertheless supplied. For each schema-selected `McpServer::Acp` entry in a session setup request, the polyfill: -1. Creates or reuses a connection-scoped localhost bridge endpoint for the - `serverId` and replaces the declaration with the HTTP transport for the - final agent. -2. Retains the native `serverId` so connections can be routed back to the - component that provided the server. -3. Opens the endpoint's native connection by sending `mcp/connect` with that - server ID toward the provider. -4. Relays requests and notifications through `mcp/message`, using the returned - `connectionId` for that active MCP connection. -5. Sends an `mcp/disconnect` request when the local transport closes and removes - the connection from the bridge. +1. Creates or reuses one connection-scoped loopback listener and replaces the + declaration with an HTTP URL whose path encodes the non-secret `serverId`. + No per-server listener or route-table entry is allocated. +2. Routes each request back to the component that owns that native registration. + The provider, not possession of the URL, decides whether it still exists. +3. Adds a runtime-only bearer credential derived from the connection secret and + server ID to the HTTP declaration's headers. The endpoint authenticates and + checks supplied Origin headers before reading the request body. Credentials + never appear in URLs; an ephemeral port alone is not access control. +4. For each POST, allocates a unique logical MCP request ID and sends + `mcp/message` to the provider. Two HTTP clients may use the same external + JSON-RPC ID without sharing routing or state. +5. Relays notifications and a final result/error for that request. Closing + its HTTP response cancels the corresponding ACP request, not the listener. Enable the polyfill crate's `unstable_session_fork` feature when adapting fork requests. Stable v1 setup includes `session/new`, `session/load`, and @@ -108,10 +115,11 @@ versions include `session/fork` when `unstable_session_fork` is enabled. Declarations using another transport are left unchanged, including extension transports represented by v2's `McpServer::Other`. -Endpoints are cached by `serverId` across session setup requests on the ACP -connection. The output declaration is rebuilt for each occurrence, preserving -that occurrence's `name`, `_meta`, and other unmodified extension fields even -when its endpoint is reused. +The same server ID derives the same route and credential on this ACP connection. +The output declaration is rebuilt for each occurrence, preserving its `name`, +`_meta`, and other unmodified extension fields. Failed setup and declaration +churn cannot accumulate per-server endpoint allocations. A server ID must never +be rebound to a different registration during the connection's lifetime. The native wire envelopes are documented in the [SDK Protocol Reference](./protocol.md#native-mcp-over-acp). @@ -119,28 +127,70 @@ Reference](./protocol.md#native-mcp-over-acp). ## HTTP Mode `McpOverAcpPolyfill::http()` is the default compatibility shape. It replaces -the native declaration with an HTTP MCP URL at `http://127.0.0.1:PORT`. The -embedded server accepts MCP POST requests and an SSE GET stream at `/`, retaining -JSON-RPC batch frames and correlating each POST with its response. +the native declaration with an HTTP MCP URL at `http://127.0.0.1:PORT/`. The +embedded server accepts a single JSON-RPC request per POST at that route, returning +JSON for a terminal-only response or SSE for a request that emits notifications. +GET and DELETE return 405. Batches and client-originated JSON-RPC responses +are rejected; there is no standalone GET event stream or MCP session ID. ```rust,ignore let bridge = McpOverAcpPolyfill::http(); ``` -The listener is bound only on loopback and uses an ephemeral port. It does not -implement resumable SSE event IDs. +Clients must send the bearer header from the declaration, both JSON and SSE +Accept types, and the required MCP protocol-version, method, and applicable +name headers. Mirrored names support MCP's Base64 sentinel encoding. Missing, +duplicate, or mismatched routing headers are rejected. -## 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 is bound only on loopback. Resumable SSE event IDs are not part of +the target MCP revision. Subscription IDs inside +`_meta["io.modelcontextprotocol/subscriptionId"]` are translated back to the +HTTP request's original ID in notifications and graceful completion results; +other metadata, progress tokens, and opaque retry state are not rewritten. -A reverse `mcp/message` request for an unknown `connectionId` receives -`Invalid params`. A reverse notification for an unknown connection is ignored, -as required for JSON-RPC notifications. +## Lifecycle and Failure Behavior -The polyfill does not infer or store ACP session IDs. Association is carried by -the declared `serverId` and the resulting active `connectionId`. +Each POST owns a pending native request, not an MCP session. Closing its response +stream cancels that request. A terminal outcome ends native work, but HTTP +admission remains held until the response body is consumed or dropped. The +listening endpoint remains available for later requests; releasing the native +registration makes requests through its old URL fail rather than reviving it. + +The adapter limits each response's queued notifications to 16 messages and +256 KiB of serialized data, admits at most 64 HTTP responses at a time, and caps +request bodies and terminal payloads at 1 MiB. The body owns the admission permit, +including while a client is not reading. A separate terminal-response path +avoids stranding completion behind a full queue. Overflow explicitly fails and +cancels that operation without blocking the shared runner or dropping events silently. + +The bridge unwraps the ACP outcome carrier before creating the HTTP JSON-RPC +response. MCP error codes/data stay MCP errors; binding failures use their +separate error codes. Queued notifications precede the terminal response. + +Unknown or late provider notifications are ignored; reverse MCP requests are +not supported. The adapter does not infer ACP session IDs or maintain MCP +initialization state. + +## Native-tool re-export contract + +The adapter creates a **new HTTP endpoint for native tool semantics**. It does +not preserve another HTTP gateway's parameter-header routing or authorization. +It removes transport-only `x-mcp-header` annotations from actual schema positions +in `tools/list` results. Argument schemas and validation keywords, tool ordering, +pagination, metadata, and similarly named properties/example/default data remain +unchanged. Annotated native tools remain listed and callable. + +Each `tools/call` issues exactly one native call, without hidden descriptor reads +or a prior client `tools/list` requirement. Native passthrough does not transform +the original descriptors. `Mcp-Param-*` headers are rejected; they confer no +authority on this endpoint. Standard MCP method/name/version header checks remain. + +If a deployment depends on an existing HTTP gateway's mirrored-parameter policy, +it must implement that policy at this endpoint or decline this re-export. + +## Validation scope + +This does not establish every optional MCP feature or complete HTTP conformance. +In particular, HTTP response limits alone do not prove native transport bounds. +Owned operation cleanup and end-to-end bounded transport are stabilization gates; +see [Native MCP-over-ACP](./mcp-over-acp.md). diff --git a/md/mcp-over-acp.md b/md/mcp-over-acp.md new file mode 100644 index 00000000..7e9c29ec --- /dev/null +++ b/md/mcp-over-acp.md @@ -0,0 +1,142 @@ +# Native MCP-over-ACP + +The native transport targets MCP 2026-07-28 only. It lets an ACP client or proxy +provide MCP tools to an agent over the existing ACP connection, without a +conductor, HTTP listener, subprocess, or MCP initialization handshake. + +Enable `unstable_mcp_over_acp` on the core SDK. Draft ACP v2 additionally +requires `unstable_protocol_v2`. The shared-schema revision is currently pinned +to a Git commit for cross-repository validation; replace that pin with the +released schema before publishing the SDK. + +## Providing tools + +Attach an `mcp_server::McpServer` to session setup through the existing builder +APIs. It publishes a `McpServer::Acp` declaration with a provider-generated +`serverId`. + +`McpService` is a reusable application service. Each `execute` call owns one +operation future and receives an `McpRequestContext` with `server_id()`, +`request_id()`, validated `metadata()`, cancellation, and an async +`send_notification` method. Share tool implementations, caches, and connection +pools deliberately; never infer a request's identity or capabilities from a +previous operation. + +Use `McpServer::new_service` for a native service, or +`new_service_with_standalone` when also exposing an independent standalone +transport. The connector-based factory remains an explicit adapter for backends +that require per-operation construction; stateless MCP does not require it. + +A connector operation is admitted as one protected task before its factory is +called. That owner drives both the backend and response forwarding, then drops +the backend and joins scoped cleanup before releasing the logical ID or replying. +If the backend exits, already-accepted output is drained without waiting for +escaped sender handles; a valid queued terminal outcome is preserved, and later +notifications are not forwarded. + +The rmcp integration's builder and `from_rmcp` use the reusable service path +for ACP attachments. Each operation uses rmcp's direct, one-request transport +without `initialize`. Its wrapper supervises rmcp handler futures through +cancellation and cleanup instead of merely dropping detached task handles. + +Custom `McpService` implementations must observe `operation_cancellation()` and +return only after their owned cleanup finishes. The binding waits for this +completion, including on ACP EOF, a runtime error, or a `connect_with` +foreground return, before dropping operation supervisors or scoped tool +runners. It does not join arbitrary user-spawned tasks or forcibly terminate +detached application work. + +The scoped `tool_fn` helpers continue to provide `McpConnectionTo` for host ACP +access. For decisions using the full MCP metadata/capabilities, implement +`McpService` or an rmcp handler receiving its `RequestContext`. Standalone MCP +connections have no ACP server or logical request ID. + +## Consuming tools + +An ACP agent holds a `ConnectionTo` or its v2 counterpart. It sends +`MessageMcpRequest::new(server_id, request_id, method)` with the inner MCP +parameters, including: + +- `io.modelcontextprotocol/protocolVersion: "2026-07-28"`; +- `io.modelcontextprotocol/clientCapabilities` as an object; +- any request-specific identity, progress token, extension settings, or retry + state required by the MCP operation. + +Choose a fresh logical request ID. It becomes the MCP JSON-RPC ID and remains +unchanged through proxies. The outer ACP request ID is separate and may change +on each hop. + +Register a `MessageMcpNotification` handler before sending requests that may +stream notifications. Route by server and logical request ID. Do not block +the ACP dispatch loop waiting for peer traffic; use a spawned task or the +connection's application future. + +The final successful ACP response is `MessageMcpResponse::Result { result, .. }` +or `MessageMcpResponse::Error { error, .. }`. Match that carrier before interpreting +the MCP outcome. The result preserves all MCP fields, including `resultType`; +the error preserves its MCP code, message, optional data, and extensions. +An MCP code must never be treated as an ACP code: for example, inner `-32000` +does not mean ACP authentication is required. + +Outer ACP failures instead describe invalid binding input, cancellation, +resource exhaustion, an unavailable registration, or a failed backend/transport. +For MRTR, process the inner `input_required` result and send a fresh request +with `inputResponses` and the exact opaque `requestState`. + +Discovery reports only the MCP revision exposed by this binding, even if the +hosted backend also supports older revisions through other transports. + +## Subscriptions and cancellation + +`subscriptions/listen` keeps one request alive. Its acknowledgement and updates +arrive as request-scoped notifications, with the logical request ID in +`io.modelcontextprotocol/subscriptionId`. An unrelated tool call does not share +that subscription's state or lifetime. + +Use `SentRequest::cancel` (or drop an unconsumed request) to cancel the outer +ACP operation. The provider revokes output immediately and stops that operation's +owned backend work; its admission slot and logical ID remain held until cleanup +finishes. Cancellation produces an outer cancellation error unless completion +already won the race. Removing a registration or receiving transport EOF cancels +its outstanding work; no separate `mcp/disconnect` exchange exists. + +## Resource limits and remaining work + +The native binding has per-registration admission and serialized payload limits. +Resource exhaustion is an outer `MCP_RESOURCE_EXHAUSTED` (`-33000`) failure, not +ACP authentication and not an inner MCP tool error. + +The transport revision introduces finite `ConnectionLimits` and `BudgetedFrame` +ownership. Adapters must keep the frame's permit through staging, deferred +dispatch, and writes; extracting a payload must not silently release its charge +while retaining the data. Async producers await capacity; synchronous dispatch +must fail explicitly instead of blocking the dispatcher needed to free capacity. +`max_queued_frames` bounds all live SDK tasks (running plus waiting), not just +waiting task slots. A persistent child connection may occupy one slot; when no +live slot remains, an ordered response callback is rejected immediately rather +than accepted behind a child that cannot finish. + +The same item-limit policy currently governs frame queues, pending requests, +live tasks, dynamic handlers, and deferred dispatch; the default is 32. +The shared payload budget defaults to 64 MiB with a 16 MiB frame maximum and +reserved response/cancellation capacity. These are serialized-payload charges, +not an exact bound on total process memory or allocations inside user code. + +Regression coverage includes sender-clone saturation, cross-budget forwarding, +retained responses and callbacks, EOF draining, and cancellation while cleanup is +paused. The [HTTP adapter](./mcp-bridge.md) separately owns its response-body permits +and fails/cancels overflowing operations. Full MCP conformance and protocol +stabilization remain separate from this implementation evidence. + +## Runnable example + +```sh +cargo run -p agent-client-protocol-rmcp \ + --example stateless_native_mcp \ + --features unstable_mcp_over_acp,unstable_protocol_v2 +``` + +This direct ACP example uses actual rmcp tools without the HTTP polyfill. +See the [protocol reference](./protocol.md#native-mcp-over-acp) for wire details +and the [RFD](https://agentclientprotocol.com/rfds/mcp-over-acp) for the design. +The [migration guide](./migration-stateless-mcp.md) lists the breaking changes. diff --git a/md/migration-rmcp-v4.md b/md/migration-rmcp-v4.md index 1de6f435..d0abc4ca 100644 --- a/md/migration-rmcp-v4.md +++ b/md/migration-rmcp-v4.md @@ -54,8 +54,9 @@ the required request metadata. ## MCP-over-ACP remains a separate draft -This upgrade does not change ACP's unstable `mcp/connect`, `mcp/message`, or -`mcp/disconnect` envelopes. Redesigning that transport around stateless, -server-addressed requests is separate work. The new transport's latest-only -target does not require removing existing rmcp behavior from this prerequisite -dependency upgrade. +The dependency upgrade alone did not change the unstable ACP wire envelopes. +The subsequent [native MCP transport migration](./migration-stateless-mcp.md) +removes the prototype's `mcp/connect`/`mcp/disconnect` lifecycle and changes +`mcp/message` to server-addressed requests with explicit outcome carriers. +Read both guides when adopting the combined major-version changes. The +latest-only native binding does not require removing standalone rmcp behavior. diff --git a/md/migration-stateless-mcp.md b/md/migration-stateless-mcp.md new file mode 100644 index 00000000..844e83ac --- /dev/null +++ b/md/migration-stateless-mcp.md @@ -0,0 +1,146 @@ +# Migrating the Native MCP Transport + +This draft replaces the connection-oriented MCP-over-ACP prototype with a +request-scoped binding for **MCP 2026-07-28 only**. It is part of the next major +SDK change, not a compatibility layer for older MCP revisions. The +`unstable_mcp_over_acp` gate remains; draft ACP v2 still has its separate gate. + +## Wire changes + +| Previous prototype | New binding | +| --- | --- | +| `mcp/connect` and `mcp/disconnect` | Removed | +| `McpConnectionId` / `connectionId` | Removed | +| `mcp/message(connectionId, method, params)` | `mcp/message(serverId, requestId, method, params)` | +| MCP initialization and connection-scoped capabilities | Required version/capabilities in each request's inner `_meta` | +| Arbitrary reverse MCP requests | MRTR `input_required` results and explicit caller retries | +| Raw MCP result or MCP error in the ACP error envelope | Successful ACP response containing exactly one inner `result` or `error` | +| HTTP MCP sessions and standalone GET streams | Independent POSTs, including long-lived subscription POSTs | + +Keep the server declaration's `serverId`. Generate a fresh logical +`McpRequestId` per call and pass it to +`MessageMcpRequest::new(server_id, request_id, method)`. That ID stays unchanged +through proxies; it is not the hop-local ACP JSON-RPC ID. + +Every inner request includes: + +```json +{ + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } +} +``` + +There is no hidden initialization or discovery prerequisite. Explicitly select +2026-07-28 when constructing an rmcp client: rmcp 3.4's default version constant +still selects an older revision. + +## Handle two error domains + +First handle the outer ACP request result, then match `MessageMcpResponse`: + +- `Result { result, .. }` contains an opaque MCP result, including any MCP + metadata, `resultType`, or explicit JSON null. +- `Error { error, .. }` contains an `McpError`. Its `code` is a plain MCP integer, + not ACP's `ErrorCode`. `data` preserves omission separately from JSON null, + and unknown error extensions survive. +- An outer ACP error reports a binding failure: invalid envelope, cancellation, + resource limit, unavailable registration, or backend/transport failure. + +Do not run ACP authentication handling on an inner MCP error code. A tool +execution failure with `isError` remains an MCP result. MRTR's `input_required` +also remains a result; retry with fresh IDs/metadata and unchanged opaque state. + +ACP v1 and v2 define independent response/error carrier types. They currently +use the same JSON representation, but may evolve separately. Use the types +for the negotiated ACP version and keep trait implementations version-specific. + +## Separate services from operations + +Use the reusable `McpService` abstraction for native providers. Per-operation +`McpRequestContext` contains logical/server identity, MCP metadata/capabilities, +cancellation, and request-scoped notification permissions. A service can share +application state without sharing MCP protocol state. + +`McpServer::new_service` registers a native service. An explicit factory/standalone +adapter remains available when constructing a backend per operation is actually +needed. `McpServer::from_rmcp` and the rmcp tool builder retain their attachment +entry points but execute ACP requests through the request-native service path. + +Do not detach tool work from its operation. Cancelling a queued call must prevent +it from starting; cancelling a running call must drop or stop its owned future +and supervise cleanup. Failure to deliver a cancelled tool's result must not +terminate the containing ACP connection. + +## Registration and cancellation + +A server ID names one registration during an ACP connection's lifetime. Do not +rebind a removed ID to another provider. Dropping the local registration rejects +future calls and cancels its active work; omitting a declaration from a later +setup request is not a new unadvertisement message. + +Use ACP request cancellation, not an MCP disconnect. Cancellation revokes +notifications immediately while cleanup retains the active ID and admission +permit. Independent calls, subscriptions, and the reusable service remain alive. +Transport EOF must begin this cleanup even if application code is still waiting +on the disconnected peer. + +## HTTP clients + +The local polyfill re-exports native tools through one signed, loopback HTTP +endpoint per ACP connection. Pass the declaration's Authorization header, never +put its bearer credential in a URL. Requests use current MCP headers and do not +exchange session IDs or `initialize`. + +The endpoint strips transport-only `x-mcp-header` schema annotations from tool +descriptors and rejects `Mcp-Param-*` headers. It does not transport an existing +HTTP gateway's routing or authorization policy. Direct tool calls require no +preliminary descriptor fetch. See the [HTTP adapter contract](./mcp-bridge.md). + +## Custom transports and connectors + +`ConnectTo::into_channel_and_future` now returns `(Channel, ConnectionDriver)`. +Wrap an owned driver future in `ConnectionDriver::new`; use +`ConnectionDriver::passive()` only for an endpoint driven elsewhere. Awaiting +the driver remains supported. Do not treat passive-driver completion as EOF: +doing so drops final responses when an input stream half-closes. + +`Channel::rx` yields `BudgetedFrame`, not a bare wire frame. Use `.frame()` to +inspect it, and preserve the envelope when forwarding through a sink. If a +custom adapter separates payload from accounting with `.into_parts()`, retain +the permit as long as the deferred payload or serialized output exists. + +For a new raw frame, use `FrameSender::send_frame(frame).await` outside dispatch +or `try_send(frame)` for explicit fail-fast admission. The old `unbounded_send` +API is removed; ignoring capacity errors silently loses protocol traffic. +Finite queue and byte policies are configured through `ConnectionLimits`. + +`RawJsonRpcMessage::Response` now carries the SDK's `RawJsonRpcResponse`, not +`schema::v1::Response`. Update raw response patterns to import +`agent_client_protocol::RawJsonRpcResponse`. Its `RawJsonRpcError` has an `i32` +code, `MaybeUndefined` data, and an extension map. Forward raw responses +unchanged to retain explicit null and unknown error fields. The error branch +stores `Box` so extensible errors do not enlarge every frame. + +`RawJsonRpcMessage::response(id, Result)` remains the convenience +constructor for ACP results. Other protocols should construct +`RawJsonRpcResponse` directly. Use `into_acp_error()` only when intentionally +dispatching an ACP error, not when forwarding MCP errors. Typed ACP request +consumers still receive the existing `Error` type. + +## Release checklist + +- Replace the draft Git schema pin with the released matching schema version. +- Coordinate major releases for crates whose public transport or rmcp-facing + API changed; do not infer compatibility solely from unchanged Cargo numbers. +- Exercise v1 and v2 carrier/error behavior, cancellation and EOF, MRTR, + subscriptions, and slow consumers before stabilizing. +- Follow the bounded transport API's ownership rules when writing custom + adapters: moving a payload must not release its accounting while a deferred + dispatch, writer, or unread HTTP body still retains it. + +The [native guide](./mcp-over-acp.md) and [protocol reference](./protocol.md#native-mcp-over-acp) +describe the target behavior. Historical migration chapters describe earlier +releases and are not a specification for this binding. diff --git a/md/protocol.md b/md/protocol.md index d30d6747..d342846b 100644 --- a/md/protocol.md +++ b/md/protocol.md @@ -11,9 +11,7 @@ unstable and is available only with the `unstable_mcp_over_acp` feature. | --- | --- | --- | | `_proxy/initialize` | request | Initialize a component as a proxy | | `_proxy/successor` | request or notification | Forward one inner ACP message to the next component | -| `mcp/connect` | request | Open a connection to an ACP-provided MCP server | -| `mcp/message` | request or notification | Carry one inner MCP message over ACP | -| `mcp/disconnect` | request | Close an MCP-over-ACP connection | +| `mcp/message` | agent request or provider notification | Invoke an MCP operation or carry a notification for that operation | There are no separate request and notification method names for successor or MCP message forwarding. The presence of an outer JSON-RPC `id` distinguishes a @@ -61,10 +59,11 @@ inner message. ## Native MCP-over-ACP -Enable `unstable_mcp_over_acp` to use the draft native transport. A component -providing an MCP server adds `McpServer::Acp` to session setup requests -(`session/new`, `session/load`, `session/resume`, and the opt-in `session/fork`). -Its wire shape contains a human-readable name and an opaque server identifier: +Enable `unstable_mcp_over_acp` to use the draft native transport targeting MCP +2026-07-28 only. ACP initialization is unchanged; there is no MCP initialization +or connect/disconnect lifecycle. A provider adds `McpServer::Acp` to session +setup requests (`session/new`, `session/resume`, v1 `session/load`, and the +opt-in `session/fork`): ```json { @@ -74,95 +73,116 @@ Its wire shape contains a human-readable name and an opaque server identifier: } ``` -`serverId` identifies the declared server and is used to route `mcp/connect` -back to the component that provided it. A provider must not reuse one server ID -for multiple visible servers on the same ACP connection. The high-level +`serverId` identifies the declared server and is used to route `mcp/message` +back to the component that provided it. A provider must not rebind a server ID +to another registration on the same ACP connection, even after removal. The high-level `agent_client_protocol::mcp_server::McpServer` APIs create this declaration automatically. -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. +An agent advertises `agentCapabilities.mcpCapabilities.acp: true` in v1 or +`capabilities.session.mcp.acp: {}` in draft v2. An optional +[HTTP adapter](./mcp-bridge.md) is only for agents with a modern MCP HTTP client. +Advertising HTTP support alone does not establish MCP revision compatibility. -### `mcp/connect` +### `mcp/message` -The MCP client opens a connection to the declared server ID: +An agent sends one request addressed to the server, with a fresh logical MCP +request ID. This ID remains unchanged through proxies even if the outer ACP +JSON-RPC ID is renumbered: ```json { "jsonrpc": "2.0", - "id": 20, - "method": "mcp/connect", - "params": { "serverId": "mcp-server:01" } + "id": 21, + "method": "mcp/message", + "params": { + "serverId": "mcp-server:01", + "requestId": "mcp-request:01", + "method": "tools/call", + "params": { + "name": "example", + "arguments": {}, + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + "progressToken": "caller-supplied-token" + } + } + } } ``` -The provider creates one active MCP connection and returns a distinct -connection ID: +The successful outer ACP response contains exactly one MCP outcome: ```json { "jsonrpc": "2.0", - "id": 20, - "result": { "connectionId": "mcp-connection:01" } + "id": 21, + "result": { + "result": { "resultType": "complete", "content": [] } + } } ``` -The server ID selects what to connect to; the connection ID selects that -particular running connection. All subsequent messages use the connection ID. +An MCP protocol error uses `{"error": {"code": ..., "message": ..., "data": ...}}` +inside the successful outer `result`, not an ACP error response. Each version's +`MessageMcpResponse::{Result, Error}` type preserves this distinction. Inner +results are opaque JSON (including null); inner error data distinguishes null +from omission. MCP error codes never acquire ACP meanings. -### `mcp/message` +Outer ACP errors describe binding failures: invalid envelope/duplicate ID +(`-32602`), cancellation (`-32800`), resource exhaustion (`-33000`), unavailable +registration (`-33001`), or backend/transport failure (`-33002`). -`mcp/message` carries one inner MCP method and its named parameters. The method -is bidirectional because MCP clients and servers can both issue requests: +MRTR `input_required` is an MCP result, not a reverse RPC; retry the original +operation with fresh metadata/IDs and unchanged opaque state. + +For `server/discover`, supported versions are restricted to the revision +exposed by this binding; a backend must actually support that revision. + +A provider may send notifications belonging to that operation: ```json { "jsonrpc": "2.0", - "id": 21, "method": "mcp/message", "params": { - "connectionId": "mcp-connection:01", - "method": "tools/call", - "params": { - "name": "example", - "arguments": {} - } + "serverId": "mcp-server:01", + "requestId": "mcp-request:01", + "method": "notifications/progress", + "params": { "progressToken": "caller-supplied-token", "progress": 1 } } } ``` -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. +Progress requires a corresponding token in the original request's inner MCP +metadata. Subscription notifications carry the listen request's logical +`requestId` in `io.modelcontextprotocol/subscriptionId`; acknowledgement comes +first. Notifications stop when their operation completes. -### `mcp/disconnect` +Both envelope types require non-null `serverId`, `requestId`, and `method` +strings. Inner `params` accepts an object or `null`; omission and `null` both +mean no parameters. A valid modern request still needs its required +`params._meta`. Optional outer ACP `_meta` is distinct from inner MCP metadata. -Disconnect is a request so the caller knows that the provider has released the -active connection: +### Cancellation and lifetime -```json -{ - "jsonrpc": "2.0", - "id": 22, - "method": "mcp/disconnect", - "params": { "connectionId": "mcp-connection:01" } -} -``` - -A successful disconnect returns an empty result: +Use [`$/cancel_request`](./request-cancellation.md) with the outer ACP request +ID. Normal proxy forwarding maps this cancellation hop by hop. It never +rewrites the logical MCP ID. Cancellation is best effort; advertising this +transport does not guarantee that every operation can be cancelled or impose +an additional cancellation support requirement. -```json -{ - "jsonrpc": "2.0", - "id": 22, - "result": {} -} -``` +Each operation owns its backend work. A result, error, cancellation, or +registration removal ends that operation; sibling requests and subscriptions stay +independent. When the SDK honors cancellation, it revokes output but keeps the +admission slot and logical ID until owned cleanup finishes. There is no MCP +connection ID to release. `server/discover` is an ordinary optional request, +not a prerequisite for tool calls. ## Related Documentation +- [Native MCP-over-ACP](./mcp-over-acp.md) - [Conductor Design](./conductor.md) - [MCP Bridge](./mcp-bridge.md) - [Original P/ACP Design Proposal](./proxying-acp.md) (historical) diff --git a/md/request-cancellation.md b/md/request-cancellation.md index 5c3a56b4..681cc162 100644 --- a/md/request-cancellation.md +++ b/md/request-cancellation.md @@ -42,6 +42,21 @@ an unknown or already-completed request ID is silently ignored. A not a string, number, or null) is logged and ignored without a reply, like any other malformed notification. +### Cancelling before publication + +`send_request` queues a request; it does not prove that the peer has received it. +If cancellation reaches the outgoing actor before it publishes that request, +the SDK settles it locally with `-32800` and sends neither the request nor its +cancellation notification. This also applies to requests waiting for session +readiness. Once published, the peer's cooperative cancellation rules apply. + +Cancellation uses a separate bounded urgent queue so it can bypass a blocked +readiness gate regardless of ordinary queue occupancy. It may therefore +overtake ordinary messages that have not yet reached the transport. Tests or +applications that need to cancel work already running on a peer must establish +that the peer has started it, rather than relying on a synchronous `send_request` +call or a scheduler yield. + ## Interoperability Protocol-level (`$/`-prefixed) notifications are optional by design. The SDK diff --git a/md/transport-architecture.md b/md/transport-architecture.md index 8204b77e..fc5a1879 100644 --- a/md/transport-architecture.md +++ b/md/transport-architecture.md @@ -29,7 +29,7 @@ The architecture separates two distinct responsibilities: This separation enables: -- **In-process efficiency**: Components in the same process can skip serialization +- **In-process efficiency**: Components can pass frames without a serialize/parse round trip - **Transport flexibility**: Easy to add new transport types (WebSockets, named pipes, etc.) - **Testability**: Mock transports for unit testing - **Clarity**: Clear boundaries between protocol and I/O concerns @@ -45,18 +45,25 @@ by the JSON-RPC envelope types from `agent-client-protocol-schema`: enum RawJsonRpcMessage { Request(Request), Notification(Notification), - Response(Response), + Response(RawJsonRpcResponse), } ``` +`RawJsonRpcResponse` uses the shared JSON-RPC response envelope with an opaque +JSON result and `RawJsonRpcError`. Raw errors keep numeric codes uninterpreted, +preserve unknown error fields, and distinguish omitted `data` from explicit +null. Only the typed ACP dispatcher converts them to ACP `Error`. Raw relays +and MCP connectors must not make that conversion: it would discard extensions +and apply the wrong protocol's error-code meaning. + At that boundary: - **Above**: Protocol layer works with application types (`OutgoingMessage`, `UntypedMessage`) - **Below**: Transport actors parse and serialize JSON-RPC frames - **Boundary**: `TransportFrame` carries one raw message, a structurally non-empty batch, or a malformed wire value retained for a relay -- **In-process API**: `Channel::rx` and `Channel::tx` carry `TransportFrame` - directly, so adapters cannot accidentally flatten a batch +- **In-process API**: `Channel::rx` and `Channel::tx` carry `BudgetedFrame` + envelopes containing the complete `TransportFrame` and its byte permit - **Failures**: I/O and connection failures are returned by the future driving a transport; they are not sent as channel entries @@ -76,8 +83,8 @@ These actors live in the protocol connection core and understand JSON-RPC semant #### Outgoing Protocol Actor ``` -Input: mpsc::UnboundedReceiver -Output: mpsc::UnboundedSender +Input: Bounded application admission queues +Output: FrameSender (BudgetedFrame) ``` Responsibilities: @@ -89,7 +96,7 @@ Responsibilities: #### Incoming Protocol Actor ``` -Input: mpsc::UnboundedReceiver +Input: FrameReceiver (BudgetedFrame) Output: Routes to pending request awaiters or registered handlers ``` @@ -115,6 +122,18 @@ The shared pending-reply registry manages request/response correlation: Runs user-spawned concurrent tasks via `cx.spawn()`. +Admission counts queued and running tasks together, using +`ConnectionLimits::max_queued_frames`. Accepted tasks are polled concurrently; +there is no second waiting pool behind permanent child connection drivers. +When all live slots are occupied, spawning or registering an ordered response +consumer fails immediately instead of accepting work that cannot make progress. +The connection's own transport driver runs outside this task pool. + +Native MCP supervisors also register protected cleanup acknowledgments. +Connection shutdown signals their cancellation and continues driving them and +their scoped tool runners until cleanup finishes, including when another task +or transport fails. Unrelated user tasks are not joined indefinitely. + ### Transport Actors These actors are driven by physical transport components. They understand @@ -124,7 +143,7 @@ correlate responses with pending requests: #### Transport Outgoing Actor ``` -Input: mpsc::UnboundedReceiver +Input: FrameReceiver (BudgetedFrame) Output: Writes to I/O (byte stream, channel, socket, etc.) ``` @@ -135,13 +154,13 @@ For byte streams: For in-process channels: -- Directly forward `TransportFrame` to the channel +- Forward `BudgetedFrame` to preserve both the frame and its admission #### Transport Incoming Actor ``` Input: Reads from I/O (byte stream, channel, socket, etc.) -Output: mpsc::UnboundedSender +Output: FrameSender (BudgetedFrame) ``` For byte streams: @@ -156,7 +175,7 @@ For byte streams: For in-process channels: -- Directly forward `TransportFrame` from the channel +- Forward `BudgetedFrame` from the channel without releasing admission The public `Channel` boundary preserves complete frames. The SDK continues to initiate requests and notifications as individual JSON-RPC messages; response @@ -219,7 +238,7 @@ Outgoing Protocol Actor | - Subscribe to replies | - Convert to RawJsonRpcMessage v - | TransportFrame (single message or batch response) + | BudgetedFrame (single message or batch response, with admission) | Transport Outgoing Actor | - Serialize (byte streams) @@ -237,7 +256,7 @@ Transport Incoming Actor | - Parse (byte streams) | - Or forward directly (channels) v - | TransportFrame (single message or incoming batch) + | BudgetedFrame (single message or incoming batch, with admission) | Incoming Protocol Actor | - Route responses → pending request awaiters @@ -262,16 +281,50 @@ Ordering](./conductor.md#routing-and-ordering). is the common component and transport abstraction. `connect_to` joins a component to its counterpart and drives the connection until completion. `into_channel_and_future` exposes the canonical low-level boundary as a -`Channel` plus the future that drives the component: +`Channel` plus an explicit connection driver: ```rust,ignore -fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<()>>); +fn into_channel_and_future(self) -> (Channel, ConnectionDriver); ``` -The returned future owns transport failures and lifecycle completion. The -channel carries only `TransportFrame` wire events. Most components implement -only `connect_to`; direct transports override `into_channel_and_future` to avoid -an intermediate copy. +`ConnectionDriver::new(future)` owns transport failures and component +completion. `ConnectionDriver::passive()` denotes an endpoint driven elsewhere, +such as an existing `Channel`; its no-op completion is **not** an EOF signal. +Both implement `Future`, so drivers can still be joined with application work. +Dynamic connectors preserve this distinction. + +A bridge must poll both copy directions while an active component runs. When +the component finishes, drain its accepted output without requiring the remote +sender to close. Between two passive endpoints, preserve independent half-close: +input EOF must still allow a final response in the other direction. + +The channel carries `BudgetedFrame` values containing complete `TransportFrame` +wire events and their resource permits. Most components implement only +`connect_to`; direct transports override `into_channel_and_future` to avoid +an intermediate copy. A forwarded frame keeps its permit through any adapter +queue, deferred dispatch, or writer. This accounting is internal and does not +change the JSON-RPC wire shape. + +### Bounded admission + +`Channel::duplex_with_limits` accepts `ConnectionLimits`. Defaults are a 16 MiB +maximum frame, a 64 MiB shared duplex serialized-payload budget, and 32 queued +frames per direction. Responses and cancellation have reserved byte capacity. +These are serialized-data and item bounds, not an exact bound on allocator +overhead or application-owned memory. + +Frame-sink clones share item capacity, including slots reserved by +`Sink::poll_ready`. `try_send` fails immediately at capacity; async sends wait +outside the dispatcher. Receiver dequeue releases the queue slot, but the +`BudgetedFrame` keeps its byte charge through deferred processing and writing. +Forward that envelope intact. Forwarding between independently budgeted +channels must also satisfy the destination's limits. + +When retaining only metadata derived from a frame, use +`FramePermit::try_reserve_metadata` with its measured serialized size. This +reserves an independent charge in every source budget; it fails immediately +rather than waiting on the payload's own reservation. Drop the original permit +once the payload is consumed, and retain the new one with the metadata. ## Transport Implementations @@ -304,13 +357,14 @@ Use cases: ### In-Process Channel For components in the same process, `Channel::duplex()` creates paired -endpoints and skips serialization entirely. Relays forward each received -`TransportFrame` without unpacking it; this preserves batch boundaries and the -original representation of malformed wire input. +endpoints without encoding and reparsing a wire message between components. +Relays forward each received `BudgetedFrame` without unpacking it; this preserves +batch boundaries, admission, and the representation of malformed wire input. +Admission still measures serialized size to enforce the shared byte budget. Benefits: -- **Zero serialization overhead**: Messages passed by value +- **No wire round trip**: Frames are passed by value, with serialized-size accounting - **Same-process efficiency**: Ideal for conductor with in-process proxies - **Explicit wire state**: No serialize/parse round trip is required, while a malformed value received from a physical transport remains an explicit frame diff --git a/src/agent-client-protocol-conductor/src/trace.rs b/src/agent-client-protocol-conductor/src/trace.rs index 62763290..b3d6117f 100644 --- a/src/agent-client-protocol-conductor/src/trace.rs +++ b/src/agent-client-protocol-conductor/src/trace.rs @@ -11,11 +11,12 @@ use std::time::Instant; use agent_client_protocol::schema::SuccessorMessage; use agent_client_protocol::schema::v1::{ - MessageMcpNotification, MessageMcpRequest, Notification as RpcNotification, - Request as RpcRequest, RequestId, Response as RpcResponse, + MessageMcpNotification, MessageMcpRequest, MessageMcpResponse, Notification as RpcNotification, + Request as RpcRequest, RequestId, }; use agent_client_protocol::{ - DynConnectTo, JsonRpcMessage, RawJsonRpcMessage, RawJsonRpcParams, Role, UntypedMessage, + DynConnectTo, JsonRpcMessage, RawJsonRpcMessage, RawJsonRpcParams, + RawJsonRpcResponse as RpcResponse, Role, UntypedMessage, }; use rustc_hash::FxHashMap; use serde::{Deserialize, Serialize}; @@ -98,6 +99,11 @@ pub struct ResponseEvent { /// True if this is an error response. pub is_error: bool, + /// Whether an error belongs to the outer ACP binding or the inner MCP peer. + /// Older trace files omit this provenance. + #[serde(default, skip_serializing_if = "Option::is_none")] + pub error_domain: Option, + /// Response result or error object. pub payload: serde_json::Value, } @@ -181,12 +187,7 @@ impl std::fmt::Debug for TraceWriter { } struct RequestDetails { - #[expect(dead_code)] protocol: Protocol, - - #[expect(dead_code)] - method: String, - request_from: ComponentIndex, request_to: ComponentIndex, } @@ -232,13 +233,13 @@ impl TraceWriter { id: serde_json::Value, method: String, session: Option, - params: serde_json::Value, + mut params: serde_json::Value, ) { + redact_http_credentials(&mut params); self.request_details.insert( id.clone(), RequestDetails { protocol, - method: method.clone(), request_from: from, request_to: to, }, @@ -261,15 +262,17 @@ impl TraceWriter { from: ComponentIndex, to: ComponentIndex, id: serde_json::Value, - is_error: bool, - payload: serde_json::Value, + error_domain: Option, + mut payload: serde_json::Value, ) { + redact_http_credentials(&mut payload); self.write_event(&TraceEvent::Response(ResponseEvent { ts: self.elapsed(), from: format!("{from:?}"), to: format!("{to:?}"), id, - is_error, + is_error: error_domain.is_some(), + error_domain, payload, })); } @@ -282,8 +285,9 @@ impl TraceWriter { to: ComponentIndex, method: impl Into, session: Option, - params: serde_json::Value, + mut params: serde_json::Value, ) { + redact_http_credentials(&mut params); self.write_event(&TraceEvent::Notification(NotificationEvent { ts: self.elapsed(), protocol, @@ -367,13 +371,13 @@ impl TraceWriter { }; let id = id_to_json(&id); if let Some(RequestDetails { - protocol: _, - method: _, + protocol, request_from, request_to, }) = self.request_details.remove(&id) { - self.response(request_to, request_from, id, is_error, payload); + let (error_domain, payload) = response_outcome(protocol, is_error, payload); + self.response(request_to, request_from, id, error_domain, payload); } } } @@ -526,6 +530,85 @@ fn params_from_transport(params: Option) -> serde_json::Value params.map_or(serde_json::Value::Null, RawJsonRpcParams::into_value) } +/// Project the logical MCP outcome without losing whether an error came from +/// the outer ACP binding. In particular, an inner -32000 is not ACP AuthRequired. +fn response_outcome( + protocol: Protocol, + outer_error: bool, + payload: serde_json::Value, +) -> (Option, serde_json::Value) { + if outer_error { + return (Some(Protocol::Acp), payload); + } + if protocol == Protocol::Mcp { + match serde_json::from_value::(payload.clone()) { + Ok(MessageMcpResponse::Result { result, .. }) => return (None, result), + Ok(MessageMcpResponse::Error { error, .. }) => { + return ( + Some(Protocol::Mcp), + serde_json::to_value(error).expect("MCP errors contain only JSON values"), + ); + } + // Retain a malformed carrier as observed, rather than invent an + // error the peer never sent. The binding validates it separately. + _ => {} + } + } + (None, payload) +} + +/// Do not persist HTTP credentials from MCP declarations or other traced payloads. +/// Only the trace's copy is modified; transport messages retain their headers. +fn redact_http_credentials(value: &mut serde_json::Value) { + fn is_credential(name: &str) -> bool { + [ + "authorization", + "proxy-authorization", + "cookie", + "set-cookie", + "x-api-key", + ] + .iter() + .any(|candidate| name.eq_ignore_ascii_case(candidate)) + } + + match value { + serde_json::Value::Object(object) => { + match object.get_mut("headers") { + Some(serde_json::Value::Array(headers)) => { + for header in headers { + if header + .get("name") + .and_then(serde_json::Value::as_str) + .is_some_and(is_credential) + && let Some(value) = header.get_mut("value") + { + *value = serde_json::Value::String("[REDACTED]".to_owned()); + } + } + } + Some(serde_json::Value::Object(headers)) => { + for (name, value) in headers { + if is_credential(name) { + *value = serde_json::Value::String("[REDACTED]".to_owned()); + } + } + } + _ => {} + } + for value in object.values_mut() { + redact_http_credentials(value); + } + } + serde_json::Value::Array(values) => { + for value in values { + redact_http_credentials(value); + } + } + _ => {} + } +} + /// A message observed going over a channel connected to `left` and `right`. /// This could be a successor message, a mcp-over-acp message, etc. #[derive(Debug)] @@ -651,16 +734,84 @@ mod tests { use agent_client_protocol::RawJsonRpcMessage; use serde_json::json; - use super::{MessageInfo, Protocol}; + use super::{MessageInfo, Protocol, ResponseEvent, redact_http_credentials, response_outcome}; + + #[test] + fn traced_mcp_outcomes_preserve_error_domain() { + let error = json!({"code":-32000,"message":"peer error","data":null,"extension":true}); + assert_eq!( + response_outcome(Protocol::Mcp, false, json!({"error":error})), + (Some(Protocol::Mcp), error.clone()) + ); + assert_eq!( + response_outcome(Protocol::Mcp, true, error.clone()), + (Some(Protocol::Acp), error) + ); + for result in [ + json!(null), + json!({"resultType":"input_required","requestState":"opaque"}), + ] { + assert_eq!( + response_outcome(Protocol::Mcp, false, json!({"result":result})), + (None, result) + ); + } + let acp_result = json!({"result": "not an MCP carrier"}); + assert_eq!( + response_outcome(Protocol::Acp, false, acp_result.clone()), + (None, acp_result) + ); + } + + #[test] + fn older_response_traces_without_error_domain_still_deserialize() { + let response: ResponseEvent = serde_json::from_value(json!({ + "ts": 0.0, "from": "Client", "to": "Agent", "id": 1, + "is_error": true, "payload": {"code": -32602, "message": "invalid"} + })) + .unwrap(); + assert!(response.is_error); + assert_eq!(response.error_domain, None); + } + + #[test] + fn trace_credentials_are_redacted_in_nested_header_shapes() { + let original = json!({ + "params": { + "mcpServers": [{ + "type": "http", + "headers": [ + {"name": "Authorization", "value": "Bearer test-token"}, + {"name": "X-Trace-Id", "value": "keep"}, + {"name": "cOoKiE", "value": "test-cookie"} + ] + }] + }, + "other": {"headers": {"X-Api-Key": "test-key", "Accept": "application/json"}} + }); + let mut traced = original.clone(); + redact_http_credentials(&mut traced); + let headers = &traced["params"]["mcpServers"][0]["headers"]; + assert_eq!(headers[0]["value"], "[REDACTED]"); + assert_eq!(headers[1]["value"], "keep"); + assert_eq!(headers[2]["value"], "[REDACTED]"); + assert_eq!(traced["other"]["headers"]["X-Api-Key"], "[REDACTED]"); + assert_eq!(traced["other"]["headers"]["Accept"], "application/json"); + assert_eq!( + original["params"]["mcpServers"][0]["headers"][0]["value"], + "Bearer test-token" + ); + } #[test] - fn tolerant_mcp_notification_params_are_traced_as_mcp() { + fn nullable_mcp_notification_params_are_traced_as_mcp() { let RawJsonRpcMessage::Notification(notification) = RawJsonRpcMessage::notification( "mcp/message".into(), json!({ - "connectionId": "connection-1", + "serverId": "server-1", + "requestId": "request-1", "method": "notifications/progress", - "params": ["invalid named params"] + "params": null }), ) .expect("notification is valid JSON-RPC") else { diff --git a/src/agent-client-protocol-conductor/tests/initialization_v2.rs b/src/agent-client-protocol-conductor/tests/initialization_v2.rs index cd9a0fdc..2741113f 100644 --- a/src/agent-client-protocol-conductor/tests/initialization_v2.rs +++ b/src/agent-client-protocol-conductor/tests/initialization_v2.rs @@ -1007,7 +1007,7 @@ async fn v2_proxy_session_helper_reissues_cancellation_for_the_downstream_hop() let (editor_out, conductor_in) = duplex(4096); let (conductor_out, editor_in) = duplex(4096); let transport = ByteStreams::new(editor_out.compat_write(), editor_in.compat()); - let client_request_id = tokio::time::timeout( + let (client_request_id, parked_id) = tokio::time::timeout( std::time::Duration::from_secs(10), Client .v2() @@ -1027,6 +1027,10 @@ async fn v2_proxy_session_helper_reissues_cancellation_for_the_downstream_hop() let pending = cx.send_request(v2::NewSessionRequest::new("/park-session")); let client_request_id = pending.id().clone(); + let parked_id = parked_id_rx + .next() + .await + .ok_or_else(|| Error::internal_error().data("parked request channel closed"))?; pending.cancel()?; let error = pending .block_task() @@ -1039,16 +1043,12 @@ async fn v2_proxy_session_helper_reissues_cancellation_for_the_downstream_hop() .block_task() .await?; assert_eq!(response.session_id, v2::SessionId::new("normal-session")); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }), ) .await .expect("v2 proxy cancellation test timed out")?; - let parked_id = tokio::time::timeout(std::time::Duration::from_secs(2), parked_id_rx.next()) - .await - .expect("agent should observe the forwarded request") - .ok_or_else(|| Error::internal_error().data("parked request channel closed"))?; assert_ne!( parked_id, client_request_id, "each proxy hop must allocate its own request ID" @@ -1120,7 +1120,7 @@ async fn v2_proxy_resume_helper_reissues_cancellation_for_the_downstream_hop() - let (editor_out, conductor_in) = duplex(4096); let (conductor_out, editor_in) = duplex(4096); let transport = ByteStreams::new(editor_out.compat_write(), editor_in.compat()); - let client_request_id = tokio::time::timeout( + let (client_request_id, parked_id) = tokio::time::timeout( std::time::Duration::from_secs(10), Client .v2() @@ -1143,6 +1143,10 @@ async fn v2_proxy_resume_helper_reissues_cancellation_for_the_downstream_hop() - "/park-session", )); let client_request_id = pending.id().clone(); + let parked_id = parked_id_rx + .next() + .await + .ok_or_else(|| Error::internal_error().data("parked request channel closed"))?; pending.cancel()?; let error = pending .block_task() @@ -1158,16 +1162,12 @@ async fn v2_proxy_resume_helper_reissues_cancellation_for_the_downstream_hop() - .block_task() .await?; assert_eq!(response, v2::ResumeSessionResponse::new()); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }), ) .await .expect("v2 resume proxy cancellation test timed out")?; - let parked_id = tokio::time::timeout(std::time::Duration::from_secs(2), parked_id_rx.next()) - .await - .expect("agent should observe the forwarded resume request") - .ok_or_else(|| Error::internal_error().data("parked request channel closed"))?; assert_ne!( parked_id, client_request_id, "each proxy hop must allocate its own request ID" diff --git a/src/agent-client-protocol-conductor/tests/mcp_cleanup_ownership.rs b/src/agent-client-protocol-conductor/tests/mcp_cleanup_ownership.rs new file mode 100644 index 00000000..a237c1bf --- /dev/null +++ b/src/agent-client-protocol-conductor/tests/mcp_cleanup_ownership.rs @@ -0,0 +1,294 @@ +#![cfg(feature = "unstable_protocol_v2")] + +//! Keep the tool runner unpolled during cancellation, while ACP still dispatches. +//! This proves that service completion alone cannot release request admission. + +use std::{ + future::Future, + pin::pin, + sync::{ + Arc, Mutex, + atomic::{AtomicBool, Ordering}, + }, + task::Poll, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, Client, ConnectionTo, Error, Responder, RunWithConnectionTo, V2ConnectionTo, + mcp_server::{ + McpConnectionTo, McpOutcome, McpRequest, McpRequestContext, McpServer, McpService, McpTool, + }, + schema::{ProtocolVersion, v2}, +}; +use futures::{FutureExt, future::BoxFuture, task::AtomicWaker}; +use schemars::JsonSchema; +use serde::{Deserialize, Serialize}; +use serde_json::json; +use tokio::sync::oneshot; + +#[derive(Deserialize, Serialize, JsonSchema)] +struct Input { + label: String, +} + +#[derive(Default)] +struct Gate { + paused: AtomicBool, + waker: AtomicWaker, +} + +impl Gate { + fn release(&self) { + self.paused.store(false, Ordering::Release); + self.waker.wake(); + } +} + +struct PausedRunner { + runner: R, + gate: Arc, +} + +impl> RunWithConnectionTo for PausedRunner { + async fn run_with_connection_to(self, cx: ConnectionTo) -> Result<(), Error> { + let mut running = pin!(self.runner.run_with_connection_to(cx)); + futures::future::poll_fn(|cx| { + self.gate.waker.register(cx.waker()); + if self.gate.paused.load(Ordering::Acquire) { + Poll::Pending + } else { + running.as_mut().poll(cx) + } + }) + .await + } +} + +struct ReleaseOnDrop(Arc); +impl Drop for ReleaseOnDrop { + fn drop(&mut self) { + self.0.release(); + } +} + +struct SignalOnDrop(Option>); +impl Drop for SignalOnDrop { + fn drop(&mut self) { + if let Some(tx) = self.0.take() { + let _sent = tx.send(()); + } + } +} + +struct ToolService { + tool: Arc, + finished: Arc>>>, +} + +impl McpService for ToolService +where + T: McpTool + 'static, +{ + fn execute( + &self, + request: McpRequest, + cx: McpRequestContext, + ) -> BoxFuture<'static, Result> { + let tool = self.tool.clone(); + let finished = self.finished.clone(); + Box::pin(async move { + let input: Input = + serde_json::from_value(request.params.expect("parameters")["arguments"].clone()) + .map_err(Error::into_internal_error)?; + let _finished = SignalOnDrop( + (input.label == "held") + .then(|| finished.lock().unwrap().take()) + .flatten(), + ); + let result = tokio::select! { + biased; + () = cx.operation_cancellation().cancelled() => Err(Error::request_cancelled()), + result = tool.call_tool(input, cx.connection().clone()) => result, + }?; + Ok(McpOutcome::Result(json!({ + "resultType": "complete", + "content": [{"type":"text", "text":result}] + }))) + }) + } +} + +async fn exercise( + tool: T, + runner: R, + started: oneshot::Receiver<()>, + dropped: oneshot::Receiver<()>, +) -> Result<(), Error> +where + T: McpTool + 'static, + R: RunWithConnectionTo + 'static, +{ + let gate = Arc::new(Gate::default()); + let (finished_tx, finished_rx) = oneshot::channel(); + let (result_tx, result_rx) = oneshot::channel(); + let state = Arc::new(Mutex::new(Some((started, dropped, finished_rx, result_tx)))); + let agent_gate = gate.clone(); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx| { + responder.respond(v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("cleanup-agent", "1"), + )) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected native declaration"); + }; + let server_id = server.server_id.clone(); + let (started, mut dropped, finished, result_tx) = + state.lock().unwrap().take().unwrap(); + let gate = agent_gate.clone(); + let work_cx = cx.clone(); + cx.spawn(async move { + let result = + async { + let _release_on_failure = ReleaseOnDrop(gate.clone()); + let request = + |label: &str| { + v2::MessageMcpRequest::new( + server_id.clone(), "same-logical-id", "tools/call", + ).params(json!({ + "name":"tool", "arguments":{"label":label}, + "_meta": { + "io.modelcontextprotocol/protocolVersion":"2026-07-28", + "io.modelcontextprotocol/clientCapabilities":{} + } + }).as_object().unwrap().clone()) + }; + let held = work_cx.send_request(request("held")); + started.await.map_err(Error::into_internal_error)?; + gate.paused.store(true, Ordering::Release); + held.cancel()?; + finished.await.map_err(Error::into_internal_error)?; + let mut response = Box::pin(held.block_task()); + + assert!( + matches!( + dropped.try_recv(), + Err(oneshot::error::TryRecvError::Empty) + ), + "runner remains paused" + ); + assert!( + response.as_mut().now_or_never().is_none(), + "cleanup precedes response" + ); + let duplicate = work_cx + .send_request(request("duplicate")) + .block_task() + .await + .expect_err("ID remains admitted during cleanup"); + assert_eq!(i32::from(duplicate.code), -32602); + + gate.release(); + let error = response.await.expect_err("cancelled operation"); + assert_eq!(i32::from(error.code), -32800); + assert!(dropped.try_recv().is_ok(), "tool dropped before reply"); + let healthy = + work_cx.send_request(request("after")).block_task().await?; + let v2::MessageMcpResponse::Result { result, .. } = healthy else { + panic!("ID reuse after cleanup should succeed"); + }; + assert_eq!(result["content"][0]["text"], "after"); + Ok::<_, Error>(()) + } + .await; + let _sent = result_tx.send(result); + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new("cleanup-session")) + }, + agent_client_protocol::on_receive_request!(), + ); + tokio::time::timeout( + Duration::from_secs(10), + Client.v2().connect_with(agent, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("cleanup-client", "1"), + )) + .block_task() + .await?; + let server = McpServer::new_service( + ToolService { + tool: Arc::new(tool), + finished: Arc::new(Mutex::new(Some(finished_tx))), + }, + "cleanup", + PausedRunner { runner, gate }, + ); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)? + .start_session() + .block_task() + .await?; + result_rx.await.map_err(Error::into_internal_error)? + }), + ) + .await + .expect("cleanup ownership regression timed out") +} + +#[tokio::test] +async fn mutable_tool_cleanup_precedes_id_release() -> Result<(), Error> { + let (started_tx, started_rx) = oneshot::channel(); + let (dropped_tx, dropped_rx) = oneshot::channel(); + let mut signals = Some((started_tx, dropped_tx)); + let (tool, runner) = agent_client_protocol::mcp_server::tool_fn_mut( + "tool", + "cleanup probe", + async move |input: Input, _cx: McpConnectionTo| { + if input.label == "held" { + let (started, dropped) = signals.take().unwrap(); + let _drop = SignalOnDrop(Some(dropped)); + let _sent = started.send(()); + std::future::pending::<()>().await; + } + Ok(input.label) + }, + agent_client_protocol::tool_fn_mut!(), + ); + exercise(tool, runner, started_rx, dropped_rx).await +} + +#[tokio::test] +async fn concurrent_tool_cleanup_precedes_id_release() -> Result<(), Error> { + let (started_tx, started_rx) = oneshot::channel(); + let (dropped_tx, dropped_rx) = oneshot::channel(); + let signals = Mutex::new(Some((started_tx, dropped_tx))); + let (tool, runner) = agent_client_protocol::mcp_server::tool_fn( + "tool", + "cleanup probe", + async move |input: Input, _cx: McpConnectionTo| { + if input.label == "held" { + let (started, dropped) = signals.lock().unwrap().take().unwrap(); + let _drop = SignalOnDrop(Some(dropped)); + let _sent = started.send(()); + std::future::pending::<()>().await; + } + Ok(input.label) + }, + agent_client_protocol::tool_fn!(), + ); + exercise(tool, runner, started_rx, dropped_rx).await +} diff --git a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs index e2adc450..e2d936d0 100644 --- a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs +++ b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill.rs @@ -6,15 +6,15 @@ use std::sync::{Arc, Mutex}; use agent_client_protocol::schema::ProtocolVersion; use agent_client_protocol::schema::v1::{ - AgentCapabilities, ConnectMcpRequest, ConnectMcpResponse, InitializeRequest, - InitializeResponse, LoadSessionRequest, LoadSessionResponse, McpCapabilities, McpServer, - McpServerAcp, NewSessionRequest, NewSessionResponse, ResumeSessionRequest, + AgentCapabilities, InitializeRequest, InitializeResponse, LoadSessionRequest, + LoadSessionResponse, McpCapabilities, McpServer, McpServerAcp, MessageMcpRequest, + MessageMcpResponse, NewSessionRequest, NewSessionResponse, ResumeSessionRequest, ResumeSessionResponse, SessionCapabilities, SessionResumeCapabilities, }; use agent_client_protocol::{Agent, Client, Conductor, ConnectTo, Proxy}; use agent_client_protocol_conductor::{ConductorImpl, ProxiesAndAgent}; use agent_client_protocol_polyfill::mcp_over_acp::McpOverAcpPolyfill; -use tokio::io::duplex; +use tokio::io::{AsyncReadExt, AsyncWriteExt, duplex}; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; const SERVER_NAME: &str = "shared-server"; @@ -56,7 +56,7 @@ struct RecordingAgent { } struct NativeMcpProvider { - connect_count: Arc, + request_count: Arc, } impl ConnectTo for NativeMcpProvider { @@ -69,10 +69,12 @@ impl ConnectTo for NativeMcpProvider { .name("native-mcp-provider") .on_receive_request_from( Agent, - async move |request: ConnectMcpRequest, responder, _cx| { + async move |request: MessageMcpRequest, responder, _cx| { assert_eq!(request.server_id.to_string(), SERVER_ID); - self.connect_count.fetch_add(1, Ordering::SeqCst); - responder.respond(ConnectMcpResponse::new("test-connection-id")) + self.request_count.fetch_add(1, Ordering::SeqCst); + responder.respond(serde_json::from_value::( + serde_json::json!({"result":{"tools": []}}), + )?) }, agent_client_protocol::on_receive_request!(), ) @@ -144,6 +146,28 @@ fn native_server() -> McpServer { McpServer::Acp(McpServerAcp::new(SERVER_NAME, SERVER_ID).meta(meta)) } +async fn http_post(url: &str, bearer: &str, id: i64) -> serde_json::Value { + let (address, route) = url + .strip_prefix("http://") + .unwrap() + .split_once('/') + .unwrap(); + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + let body = serde_json::json!({"jsonrpc":"2.0","id":id,"method":"tools/list", + "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28", + "io.modelcontextprotocol/clientCapabilities":{}}}}) + .to_string(); + let request = format!( + "POST /{route} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: tools/list\r\nContent-Length: {}\r\n\r\n{body}", + body.len() + ); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut response = String::new(); + stream.read_to_string(&mut response).await.unwrap(); + assert!(response.starts_with("HTTP/1.1 200"), "{response}"); + serde_json::from_str(response.split("\r\n\r\n").nth(1).unwrap()).unwrap() +} + async fn recv( response: agent_client_protocol::SentRequest, ) -> Result { @@ -158,7 +182,7 @@ async fn recv( async fn run_with_polyfill( agent: RecordingAgent, - provider_connect_count: Arc, + provider_request_count: Arc, editor_task: impl AsyncFnOnce( agent_client_protocol::ConnectionTo, ) -> Result<(), agent_client_protocol::Error>, @@ -184,7 +208,7 @@ async fn run_with_polyfill( "polyfill-test-conductor".to_string(), ProxiesAndAgent::new(agent) .proxy(NativeMcpProvider { - connect_count: provider_connect_count, + request_count: provider_request_count, }) .proxy(McpOverAcpPolyfill::http()), ) @@ -206,9 +230,9 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ capabilities: agent_capabilities(McpCapabilities::new().http(true)), observed: observed.clone(), }; - let connect_count = Arc::new(AtomicUsize::new(0)); + let request_count = Arc::new(AtomicUsize::new(0)); - run_with_polyfill(agent, connect_count.clone(), async |connection| { + run_with_polyfill(agent, request_count.clone(), async |connection| { let initialize = recv(connection.send_request(InitializeRequest::new(ProtocolVersion::V1))).await?; assert!(initialize.agent_capabilities.mcp_capabilities.http); @@ -235,6 +259,20 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ )) .await?; + let (url, bearer) = { + let setup = observed.setup.lock().unwrap(); + let McpServer::Http(server) = &setup[0].mcp_servers[0] else { + panic!("expected HTTP declaration") + }; + (server.url.clone(), server.headers[0].value.clone()) + }; + let (first, second) = + tokio::join!(http_post(&url, &bearer, 1), http_post(&url, &bearer, 1),); + assert_eq!( + first, + serde_json::json!({"jsonrpc":"2.0","id":1,"result":{"tools":[]}}) + ); + assert_eq!(second, first); Ok(()) }) .await?; @@ -244,9 +282,9 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ .lock() .expect("setup request mutex should not be poisoned"); assert_eq!( - connect_count.load(Ordering::SeqCst), - 1, - "one reused listener should create one native MCP connection" + request_count.load(Ordering::SeqCst), + 2, + "each HTTP POST creates exactly one native MCP request, without a connect handshake" ); assert_eq!(setup.len(), 3); assert_eq!(setup[0].method, SetupMethod::New); @@ -267,7 +305,9 @@ async fn http_downstream_receives_stable_transformed_declarations_for_all_setup_ }; assert_eq!(server.name, SERVER_NAME); assert_eq!(server.meta.as_ref(), Some(&expected_meta)); - assert!(server.headers.is_empty()); + assert_eq!(server.headers.len(), 1); + assert_eq!(server.headers[0].name, "Authorization"); + assert!(server.headers[0].value.starts_with("Bearer ")); assert!(server.url.starts_with("http://127.0.0.1:")); if let Some(endpoint) = &endpoint { assert_eq!( @@ -292,9 +332,9 @@ async fn native_downstream_keeps_capability_and_declaration_unchanged() }; let declaration = native_server(); let expected = declaration.clone(); - let connect_count = Arc::new(AtomicUsize::new(0)); + let request_count = Arc::new(AtomicUsize::new(0)); - run_with_polyfill(agent, connect_count.clone(), async move |connection| { + run_with_polyfill(agent, request_count.clone(), async move |connection| { let initialize = recv(connection.send_request(InitializeRequest::new(ProtocolVersion::V1))).await?; assert!(!initialize.agent_capabilities.mcp_capabilities.http); @@ -315,7 +355,7 @@ async fn native_downstream_keeps_capability_and_declaration_unchanged() assert_eq!(setup.len(), 1); assert_eq!(setup[0].mcp_servers, vec![expected]); assert_eq!( - connect_count.load(Ordering::SeqCst), + request_count.load(Ordering::SeqCst), 0, "a native-capable downstream should not be routed through the HTTP adapter" ); diff --git a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs index 284ee0ce..d28197a0 100644 --- a/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs +++ b/src/agent-client-protocol-conductor/tests/mcp_over_acp_polyfill_v2.rs @@ -1,9 +1,8 @@ #![cfg(feature = "unstable_protocol_v2")] -//! V2 integration coverage for the public MCP-over-ACP compatibility proxy. +//! End-to-end v2 coverage for the request-scoped MCP HTTP adapter. use std::{ - collections::BTreeMap, path::PathBuf, sync::{ Arc, Mutex, @@ -17,71 +16,32 @@ use agent_client_protocol::{ }; use agent_client_protocol_conductor::{ConductorImpl, ProxiesAndAgent}; use agent_client_protocol_polyfill::mcp_over_acp::McpOverAcpPolyfill; -use rmcp::{ - ServiceExt as _, - transport::{ - StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig, - }, -}; -use tokio::io::duplex; +use tokio::io::{AsyncReadExt, AsyncWriteExt, duplex}; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; -const SERVER_NAME: &str = "shared-v2-server"; -const SERVER_ID: &str = "shared-v2-server-id"; +const SERVER_ID: &str = "v2-server-id"; -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -enum SetupMethod { - New, - Resume, -} - -#[derive(Debug)] -struct SetupRequest { - method: SetupMethod, - mcp_servers: Vec, -} - -#[derive(Default)] -struct ObservedRequests { - setup: Mutex>, -} - -impl ObservedRequests { - fn record(&self, method: SetupMethod, mcp_servers: Vec) { - self.setup - .lock() - .expect("setup request mutex should not be poisoned") - .push(SetupRequest { - method, - mcp_servers, - }); - } -} - -struct RecordingAgent { +struct TestAgent { capabilities: v2::AgentCapabilities, - observed: Arc, + observed: Arc>>, } -impl ConnectTo for RecordingAgent { +impl ConnectTo for TestAgent { async fn connect_to( self, client: impl ConnectTo, ) -> Result<(), agent_client_protocol::Error> { let capabilities = self.capabilities; - let new_observed = Arc::clone(&self.observed); - let resume_observed = self.observed; - + let observed = self.observed; Agent .v2() - .name("recording-v2-agent") + .name("v2-http-test-agent") .on_receive_request( async move |request: v2::InitializeRequest, responder, _cx| { - assert_eq!(request.protocol_version, ProtocolVersion::V2); responder.respond( v2::InitializeResponse::new( request.protocol_version, - implementation("recording-v2-agent"), + v2::Implementation::new("test", "1.0.0"), ) .capabilities(capabilities.clone()), ) @@ -90,15 +50,8 @@ impl ConnectTo for RecordingAgent { ) .on_receive_request( async move |request: v2::NewSessionRequest, responder, _cx| { - new_observed.record(SetupMethod::New, request.mcp_servers); - responder.respond(v2::NewSessionResponse::new("v2-session-id")) - }, - agent_client_protocol::on_receive_request!(), - ) - .on_receive_request( - async move |request: v2::ResumeSessionRequest, responder, _cx| { - resume_observed.record(SetupMethod::Resume, request.mcp_servers); - responder.respond(v2::ResumeSessionResponse::new()) + *observed.lock().unwrap() = request.mcp_servers; + responder.respond(v2::NewSessionResponse::new("session")) }, agent_client_protocol::on_receive_request!(), ) @@ -107,85 +60,82 @@ impl ConnectTo for RecordingAgent { } } -struct NativeMcpProvider { - connect_count: Arc, - request_methods: Arc>>, - notification_methods: Arc>>, - disconnect_count: Arc, -} +struct TestProvider(Arc>>, Arc, Arc); -impl ConnectTo for NativeMcpProvider { +impl ConnectTo for TestProvider { async fn connect_to( self, client: impl ConnectTo, ) -> Result<(), agent_client_protocol::Error> { - let request_methods = Arc::clone(&self.request_methods); - let notification_methods = Arc::clone(&self.notification_methods); - let disconnect_count = Arc::clone(&self.disconnect_count); - Proxy .v2() - .name("native-v2-mcp-provider") + .name("v2-mcp-provider") .on_receive_request_from( Agent, - async move |request: v2::ConnectMcpRequest, responder, _cx| { + async move |request: v2::MessageMcpRequest, responder, cx| { assert_eq!(request.server_id.to_string(), SERVER_ID); - self.connect_count.fetch_add(1, Ordering::SeqCst); - responder.respond(v2::ConnectMcpResponse::new("v2-test-connection-id")) - }, - agent_client_protocol::on_receive_request!(), - ) - .on_receive_request_from( - Agent, - async move |request: v2::MessageMcpRequest, responder, _cx| { - request_methods - .lock() - .expect("request method mutex should not be poisoned") - .push(request.method.clone()); - match request.method.as_str() { - "initialize" => { - let protocol_version = request - .params - .as_ref() - .and_then(|params| params.get("protocolVersion")) - .cloned() - .unwrap_or_else(|| serde_json::json!("2025-06-18")); - responder.respond(serde_json::from_value(serde_json::json!({ - "protocolVersion": protocol_version, - "capabilities": { - "tools": {} - }, - "serverInfo": { - "name": "v2-polyfill-test-mcp-server", - "version": env!("CARGO_PKG_VERSION") - } - }))?) - } - "tools/list" => responder - .respond(serde_json::from_value(serde_json::json!({ "tools": [] }))?), - method => responder.respond_with_error( - agent_client_protocol::Error::method_not_found().data(method), - ), + self.0.lock().unwrap().push(request.request_id.to_string()); + self.1.fetch_add(1, Ordering::SeqCst); + if request.method == "subscriptions/listen" + || request.method == "subscriptions/flood" + { + let params = if request.method == "subscriptions/flood" { + serde_json::Map::from_iter([( + "payload".into(), + serde_json::json!("x".repeat(300 * 1024)), + )]) + } else { + serde_json::Map::from_iter([( + "_meta".into(), + serde_json::json!({"io.modelcontextprotocol/subscriptionId": + request.request_id.to_string()}), + )]) + }; + cx.send_notification_to( + Agent, + v2::MessageMcpNotification::new( + SERVER_ID, + request.request_id, + "notifications/subscriptions/acknowledged", + ) + .params(params), + )?; + let cancelled = responder.cancellation(); + let count = self.2.clone(); + cx.spawn(async move { + cancelled.cancelled().await; + count.fetch_add(1, Ordering::SeqCst); + responder.respond_with_error( + agent_client_protocol::Error::request_cancelled(), + ) + })?; + return Ok(()); } - }, - agent_client_protocol::on_receive_request!(), - ) - .on_receive_notification_from( - Agent, - async move |notification: v2::MessageMcpNotification, _cx| { - notification_methods - .lock() - .expect("notification method mutex should not be poisoned") - .push(notification.method); - Ok(()) - }, - agent_client_protocol::on_receive_notification!(), - ) - .on_receive_request_from( - Agent, - async move |_request: v2::DisconnectMcpRequest, responder, _cx| { - disconnect_count.fetch_add(1, Ordering::SeqCst); - responder.respond(v2::DisconnectMcpResponse::new()) + let result = match request.method.as_str() { + "tools/list" => serde_json::json!({"tools":[ + {"name":"ping","inputSchema":{"type":"object","properties":{}}}, + {"name":"restricted","inputSchema":{"type":"object","properties":{ + "region":{"type":"string","x-mcp-header":"Region"} + }}} + ]}), + "tools/call" => serde_json::json!({"content":[]}), + "tools/error" => { + return responder.respond(serde_json::from_value::< + v2::MessageMcpResponse, + >( + serde_json::json!({"error":{"code":-32000,"message":"peer-owned", + "data":{"source":"backend"}}}), + )?); + } + _ => { + return responder.respond_with_error( + agent_client_protocol::Error::method_not_found(), + ); + } + }; + responder.respond(serde_json::from_value::( + serde_json::json!({"result":result}), + )?) }, agent_client_protocol::on_receive_request!(), ) @@ -194,78 +144,30 @@ impl ConnectTo for NativeMcpProvider { } } -fn implementation(name: &str) -> v2::Implementation { - v2::Implementation::new(name, env!("CARGO_PKG_VERSION")) -} - -fn agent_capabilities(mcp: v2::McpCapabilities) -> v2::AgentCapabilities { - v2::AgentCapabilities::new().session(v2::SessionCapabilities::new().mcp(mcp)) -} - -fn initialize_request() -> v2::InitializeRequest { - v2::InitializeRequest::new( - ProtocolVersion::V2, - implementation("v2-polyfill-test-client"), - ) -} - -fn server_meta() -> v2::Meta { - let mut meta = v2::Meta::new(); - meta.insert( - "source".to_owned(), - serde_json::Value::String("v2-integration-test".to_owned()), - ); - meta -} - -fn native_server() -> v2::McpServer { - v2::McpServer::Acp(v2::McpServerAcp::new(SERVER_NAME, SERVER_ID).meta(server_meta())) -} - -fn future_server() -> v2::McpServer { - v2::McpServer::Other(v2::OtherMcpServer::new( - "_future_transport", - BTreeMap::from([ - ("name".to_owned(), serde_json::json!("future-v2-server")), - ( - "configuration".to_owned(), - serde_json::json!({ "preserve": true }), - ), - ]), - )) -} - -fn test_servers() -> Vec { - vec![native_server(), future_server()] -} - -async fn run_with_polyfill( - agent: RecordingAgent, - provider_connect_count: Arc, - provider_request_methods: Arc>>, - provider_notification_methods: Arc>>, - provider_disconnect_count: Arc, - editor_task: impl AsyncFnOnce(V2ConnectionTo) -> Result<(), agent_client_protocol::Error>, +async fn run( + capabilities: v2::AgentCapabilities, + observed: Arc>>, + ids: Arc>>, + count: Arc, + cancelled: Arc, + editor: impl AsyncFnOnce(V2ConnectionTo) -> Result<(), agent_client_protocol::Error>, ) -> Result<(), agent_client_protocol::Error> { let (editor_out, conductor_in) = duplex(4096); let (conductor_out, editor_in) = duplex(4096); let transport = agent_client_protocol::ByteStreams::new(editor_out.compat_write(), editor_in.compat()); - Client .v2() - .name("v2-polyfill-test-client") + .name("v2-mcp-test-client") .with_spawned(|_cx| async move { ConductorImpl::new_agent( - "v2-polyfill-test-conductor", - ProxiesAndAgent::new(agent) - .proxy(NativeMcpProvider { - connect_count: provider_connect_count, - request_methods: provider_request_methods, - notification_methods: provider_notification_methods, - disconnect_count: provider_disconnect_count, - }) - .proxy(McpOverAcpPolyfill::http()), + "v2-mcp-test-conductor", + ProxiesAndAgent::new(TestAgent { + capabilities, + observed, + }) + .proxy(TestProvider(ids, count, cancelled)) + .proxy(McpOverAcpPolyfill::http()), ) .run(agent_client_protocol::ByteStreams::new( conductor_out.compat_write(), @@ -273,188 +175,176 @@ async fn run_with_polyfill( )) .await }) - .connect_with(transport, editor_task) + .connect_with(transport, editor) .await } -fn negotiated_mcp_capabilities(response: &v2::InitializeResponse) -> &v2::McpCapabilities { - response - .capabilities - .session - .as_ref() - .expect("the test agent should advertise session support") - .mcp - .as_ref() - .expect("the test agent should advertise MCP support") +fn native_server() -> v2::McpServer { + let mut meta = v2::Meta::new(); + meta.insert("preserve".into(), serde_json::json!(true)); + v2::McpServer::Acp(v2::McpServerAcp::new("native", SERVER_ID).meta(meta)) } -#[tokio::test] -async fn http_downstream_adapts_v2_capabilities_and_only_transforms_native_servers() --> Result<(), agent_client_protocol::Error> { - let observed = Arc::new(ObservedRequests::default()); - let agent = RecordingAgent { - capabilities: agent_capabilities( - v2::McpCapabilities::new().http(v2::McpHttpCapabilities::new()), - ), - observed: Arc::clone(&observed), +fn initialize() -> v2::InitializeRequest { + v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("test", "1.0.0"), + ) +} + +async fn post(url: &str, bearer: &str, method: &str, tool: &str) -> serde_json::Value { + let (address, route) = url + .strip_prefix("http://") + .unwrap() + .split_once('/') + .unwrap(); + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + let mut params = serde_json::json!({ + "_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28", + "io.modelcontextprotocol/clientCapabilities":{}} + }); + if method == "tools/call" { + params["name"] = serde_json::json!(tool); + params["arguments"] = serde_json::json!({}); + } + let body = serde_json::json!({"jsonrpc":"2.0","id":"same","method":method, + "params":params}) + .to_string(); + let name = if method == "tools/call" { + format!("Mcp-Name: {tool}\r\n") + } else { + String::new() }; - let connect_count = Arc::new(AtomicUsize::new(0)); - let request_methods = Arc::new(Mutex::new(Vec::new())); - let notification_methods = Arc::new(Mutex::new(Vec::new())); - let disconnect_count = Arc::new(AtomicUsize::new(0)); + let request = format!( + "POST /{route} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: {method}\r\n{name}Content-Length: {}\r\n\r\n{body}", + body.len() + ); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut response = String::new(); + stream.read_to_string(&mut response).await.unwrap(); + assert!(response.starts_with("HTTP/1.1 200"), "{response}"); + serde_json::from_str(response.split("\r\n\r\n").nth(1).unwrap()).unwrap() +} - run_with_polyfill( - agent, - Arc::clone(&connect_count), - Arc::clone(&request_methods), - Arc::clone(¬ification_methods), - Arc::clone(&disconnect_count), +#[tokio::test] +async fn modern_http_v2_requests_are_stateless_and_isolated() +-> Result<(), agent_client_protocol::Error> { + let observed = Arc::new(Mutex::new(Vec::new())); + let ids = Arc::new(Mutex::new(Vec::new())); + let count = Arc::new(AtomicUsize::new(0)); + let caps = v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new() + .mcp(v2::McpCapabilities::new().http(v2::McpHttpCapabilities::new())), + ); + run( + caps, + observed.clone(), + ids.clone(), + count.clone(), + Arc::default(), async |connection| { - let initialize = connection - .send_request(initialize_request()) - .block_task() - .await?; - let mcp = negotiated_mcp_capabilities(&initialize); - assert!(mcp.http.is_some()); + let initialized = connection.send_request(initialize()).block_task().await?; assert!( - mcp.acp.is_some(), - "the HTTP adapter should advertise v2 native MCP support upstream" + initialized + .capabilities + .session + .unwrap() + .mcp + .unwrap() + .acp + .is_some() ); - - let cwd = PathBuf::from("/tmp"); - let session = connection - .send_request(v2::NewSessionRequest::new(cwd.clone()).mcp_servers(test_servers())) - .block_task() - .await?; connection .send_request( - v2::ResumeSessionRequest::new(session.session_id, cwd) - .mcp_servers(test_servers()), + v2::NewSessionRequest::new(PathBuf::from("/tmp")) + .mcp_servers(vec![native_server()]), ) .block_task() .await?; - - let endpoint = { - let setup = observed - .setup - .lock() - .expect("setup request mutex should not be poisoned"); - let v2::McpServer::Http(server) = &setup[0].mcp_servers[0] else { - panic!("expected the native declaration to be adapted to HTTP") + let (url, bearer) = { + let observed = observed.lock().unwrap(); + let v2::McpServer::Http(server) = &observed[0] else { + panic!("expected HTTP endpoint") }; - server.url.clone() + assert_eq!( + server.meta.as_ref().unwrap().get("preserve"), + Some(&serde_json::json!(true)) + ); + assert_eq!(server.headers[0].name, "Authorization"); + (server.url.clone(), server.headers[0].value.clone()) }; - let mcp_client = () - .serve(StreamableHttpClientTransport::from_config( - StreamableHttpClientTransportConfig::with_uri(endpoint), - )) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - let tools = mcp_client - .list_tools(None) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - assert!(tools.tools.is_empty()); - mcp_client - .cancel() - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; + // A direct call does not require discovery or an internal tools/list lookup. + let direct = post(&url, &bearer, "tools/call", "ping").await; + assert_eq!( + direct, + serde_json::json!({"jsonrpc":"2.0","id":"same","result":{"content":[]}}) + ); + let annotated = post(&url, &bearer, "tools/call", "restricted").await; + assert_eq!(annotated["result"], serde_json::json!({"content":[]})); + let peer_error = post(&url, &bearer, "tools/error", "").await; + assert_eq!( + peer_error["error"], + serde_json::json!({"code":-32000, + "message":"peer-owned","data":{"source":"backend"}}) + ); + let (a, b) = tokio::join!( + post(&url, &bearer, "tools/list", ""), + post(&url, &bearer, "tools/list", "") + ); + assert_eq!( + a, + serde_json::json!({"jsonrpc":"2.0","id":"same","result":{"tools":[ + {"name":"ping","inputSchema":{"type":"object","properties":{}}}, + {"name":"restricted","inputSchema":{"type":"object","properties":{ + "region":{"type":"string"} + }}} + ]}}) + ); + assert_eq!(a, b); Ok(()) }, ) .await?; - - let setup = observed - .setup - .lock() - .expect("setup request mutex should not be poisoned"); - assert_eq!( - connect_count.load(Ordering::SeqCst), - 1, - "one reused listener should create one v2 native MCP connection" - ); - assert_eq!( - *request_methods - .lock() - .expect("request method mutex should not be poisoned"), - ["initialize", "tools/list"] + assert_eq!(count.load(Ordering::SeqCst), 5); + let ids = ids.lock().unwrap(); + assert_eq!(ids.len(), 5); + assert_ne!( + ids[0], ids[1], + "external JSON-RPC IDs must not collide at the ACP hop" ); - assert_eq!( - *notification_methods - .lock() - .expect("notification method mutex should not be poisoned"), - ["notifications/initialized"] - ); - assert_eq!(setup.len(), 2); - assert_eq!(setup[0].method, SetupMethod::New); - assert_eq!(setup[1].method, SetupMethod::Resume); - - let expected_future_server = future_server(); - let expected_meta = server_meta(); - let mut endpoint = None; - for request in setup.iter() { - assert_eq!(request.mcp_servers.len(), 2); - let v2::McpServer::Http(server) = &request.mcp_servers[0] else { - panic!( - "expected the ACP declaration to become HTTP for {:?}, got {:?}", - request.method, request.mcp_servers - ); - }; - assert_eq!(server.name, SERVER_NAME); - assert_eq!(server.meta.as_ref(), Some(&expected_meta)); - assert!(server.headers.is_empty()); - assert!(server.url.starts_with("http://127.0.0.1:")); - assert_eq!( - request.mcp_servers[1], expected_future_server, - "the polyfill must preserve custom v2 MCP transports" - ); - if let Some(endpoint) = &endpoint { - assert_eq!( - &server.url, endpoint, - "the same ACP server ID should reuse one listener" - ); - } else { - endpoint = Some(server.url.clone()); - } - } - Ok(()) } #[tokio::test] -async fn native_v2_downstream_keeps_capability_and_declarations_unchanged() +async fn native_v2_declarations_pass_through_without_http() -> Result<(), agent_client_protocol::Error> { - let observed = Arc::new(ObservedRequests::default()); - let agent = RecordingAgent { - capabilities: agent_capabilities( - v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), - ), - observed: Arc::clone(&observed), - }; - let expected = test_servers(); - let connect_count = Arc::new(AtomicUsize::new(0)); - let request_methods = Arc::new(Mutex::new(Vec::new())); - let notification_methods = Arc::new(Mutex::new(Vec::new())); - let disconnect_count = Arc::new(AtomicUsize::new(0)); - - run_with_polyfill( - agent, - Arc::clone(&connect_count), - Arc::clone(&request_methods), - Arc::clone(¬ification_methods), - Arc::clone(&disconnect_count), - async move |connection| { - let initialize = connection - .send_request(initialize_request()) - .block_task() - .await?; - let mcp = negotiated_mcp_capabilities(&initialize); - assert!(mcp.http.is_none()); - assert!(mcp.acp.is_some()); - + let observed = Arc::new(Mutex::new(Vec::new())); + let caps = v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new() + .mcp(v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new())), + ); + run( + caps, + observed.clone(), + Arc::default(), + Arc::default(), + Arc::default(), + async |connection| { + let initialized = connection.send_request(initialize()).block_task().await?; + assert!( + initialized + .capabilities + .session + .unwrap() + .mcp + .unwrap() + .acp + .is_some() + ); connection .send_request( - v2::NewSessionRequest::new(PathBuf::from("/tmp")).mcp_servers(expected.clone()), + v2::NewSessionRequest::new(PathBuf::from("/tmp")) + .mcp_servers(vec![native_server()]), ) .block_task() .await?; @@ -462,105 +352,91 @@ async fn native_v2_downstream_keeps_capability_and_declarations_unchanged() }, ) .await?; - - let setup = observed - .setup - .lock() - .expect("setup request mutex should not be poisoned"); - assert_eq!(setup.len(), 1); - assert_eq!(setup[0].mcp_servers, test_servers()); - assert_eq!( - connect_count.load(Ordering::SeqCst), - 0, - "a native-capable v2 downstream should bypass the HTTP adapter" - ); - assert!( - request_methods - .lock() - .expect("request method mutex should not be poisoned") - .is_empty() - ); - assert!( - notification_methods - .lock() - .expect("notification method mutex should not be poisoned") - .is_empty() - ); - assert_eq!(disconnect_count.load(Ordering::SeqCst), 0); - + assert_eq!(*observed.lock().unwrap(), vec![native_server()]); Ok(()) } #[tokio::test] -async fn unavailable_v2_downstream_rejects_native_declarations() +async fn closing_subscription_stream_cancels_only_its_native_request() -> Result<(), agent_client_protocol::Error> { - let observed = Arc::new(ObservedRequests::default()); - let agent = RecordingAgent { - capabilities: agent_capabilities(v2::McpCapabilities::new()), - observed: Arc::clone(&observed), - }; - let connect_count = Arc::new(AtomicUsize::new(0)); - let request_methods = Arc::new(Mutex::new(Vec::new())); - let notification_methods = Arc::new(Mutex::new(Vec::new())); - let disconnect_count = Arc::new(AtomicUsize::new(0)); - - run_with_polyfill( - agent, - Arc::clone(&connect_count), - Arc::clone(&request_methods), - Arc::clone(¬ification_methods), - Arc::clone(&disconnect_count), - async move |connection| { - let initialize = connection - .send_request(initialize_request()) - .block_task() - .await?; - let mcp = negotiated_mcp_capabilities(&initialize); - assert!(mcp.http.is_none()); - assert!(mcp.acp.is_none()); - - let error = connection - .send_request( - v2::NewSessionRequest::new(PathBuf::from("/tmp")) - .mcp_servers(vec![native_server()]), - ) - .block_task() - .await - .expect_err("native MCP should require a downstream transport"); - assert_eq!(error.code, agent_client_protocol::ErrorCode::InvalidParams); - assert_eq!( - error.data, - Some(serde_json::json!( - "the downstream agent supports neither native nor HTTP MCP transport" - )) + let observed = Arc::new(Mutex::new(Vec::new())); + let cancelled = Arc::new(AtomicUsize::new(0)); + let ids = Arc::new(Mutex::new(Vec::new())); + let caps = v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new() + .mcp(v2::McpCapabilities::new().http(v2::McpHttpCapabilities::new())), + ); + run( + caps, observed.clone(), ids.clone(), Arc::default(), cancelled.clone(), + async |connection| { + connection.send_request(initialize()).block_task().await?; + connection.send_request(v2::NewSessionRequest::new(PathBuf::from("/tmp")) + .mcp_servers(vec![native_server()])).block_task().await?; + let (url, bearer) = { + let observed = observed.lock().unwrap(); + let v2::McpServer::Http(server) = &observed[0] else { panic!("expected HTTP endpoint") }; + (server.url.clone(), server.headers[0].value.clone()) + }; + let (address, route) = url.strip_prefix("http://").unwrap().split_once('/').unwrap(); + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + let body = serde_json::json!({ + "jsonrpc":"2.0","id":73,"method":"subscriptions/listen", + "params":{"notifications":{"toolsListChanged":true}, "_meta":{"io.modelcontextprotocol/protocolVersion":"2026-07-28", + "io.modelcontextprotocol/clientCapabilities":{}}} + }).to_string(); + let request = format!( + "POST /{route} HTTP/1.1\r\nHost: {address}\r\nAuthorization: {bearer}\r\nAccept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: subscriptions/listen\r\nContent-Length: {}\r\n\r\n{body}", + body.len() ); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut output = String::new(); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while !output.contains("notifications/subscriptions/acknowledged") { + let mut buf = [0; 2048]; + let n = stream.read(&mut buf).await.unwrap(); + assert_ne!(n, 0, "subscription stream closed before ack"); + output.push_str(std::str::from_utf8(&buf[..n]).unwrap()); + } + }).await.expect("expected subscription acknowledgment"); + assert!(output.contains("\"io.modelcontextprotocol/subscriptionId\":73"), "{output}"); + // A second live POST uses the same external ID, but must retain its + // own generated logical ID and cancellation lifetime. + let mut second = tokio::net::TcpStream::connect(address).await.unwrap(); + second.write_all(request.as_bytes()).await.unwrap(); + let mut second_output = String::new(); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while !second_output.contains("notifications/subscriptions/acknowledged") { + let mut buf = [0; 2048]; + let n = second.read(&mut buf).await.unwrap(); + assert_ne!(n, 0, "second subscription closed before ack"); + second_output.push_str(std::str::from_utf8(&buf[..n]).unwrap()); + } + }).await.expect("expected second subscription acknowledgment"); + assert!(second_output.contains("\"io.modelcontextprotocol/subscriptionId\":73"), "{second_output}"); + let logical_ids = ids.lock().unwrap().clone(); + assert_eq!(logical_ids.len(), 2); + assert_ne!(logical_ids[0], logical_ids[1]); + drop(stream); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while cancelled.load(Ordering::SeqCst) != 1 { + tokio::task::yield_now().await; + } + }).await.expect("closing the HTTP stream must cancel native ACP request"); + drop(second); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while cancelled.load(Ordering::SeqCst) != 2 { + tokio::task::yield_now().await; + } + }).await.expect("closing the second stream must cancel its own ACP request"); + let overflow = post(&url, &bearer, "subscriptions/flood", "").await; + assert_eq!(overflow["error"]["code"], -33000, "{overflow}"); + tokio::time::timeout(std::time::Duration::from_secs(3), async { + while cancelled.load(Ordering::SeqCst) != 3 { + tokio::task::yield_now().await; + } + }).await.expect("overflow must cancel only its native ACP request"); + assert_eq!(post(&url, &bearer, "tools/list", "").await["result"]["tools"][0]["name"], "ping"); Ok(()) }, - ) - .await?; - - assert!( - observed - .setup - .lock() - .expect("setup request mutex should not be poisoned") - .is_empty(), - "the rejected request must not reach the downstream agent" - ); - assert_eq!(connect_count.load(Ordering::SeqCst), 0); - assert!( - request_methods - .lock() - .expect("request method mutex should not be poisoned") - .is_empty() - ); - assert!( - notification_methods - .lock() - .expect("notification method mutex should not be poisoned") - .is_empty() - ); - assert_eq!(disconnect_count.load(Ordering::SeqCst), 0); - - Ok(()) + ).await } diff --git a/src/agent-client-protocol-conductor/tests/mcp_server_handler_chain_v2.rs b/src/agent-client-protocol-conductor/tests/mcp_server_handler_chain_v2.rs index 75da7d16..753130df 100644 --- a/src/agent-client-protocol-conductor/tests/mcp_server_handler_chain_v2.rs +++ b/src/agent-client-protocol-conductor/tests/mcp_server_handler_chain_v2.rs @@ -10,8 +10,8 @@ use std::{ }; use agent_client_protocol::{ - Agent, ByteStreams, Client, Conductor, ConnectTo, DynConnectTo, Error, NullRun, Proxy, - Responder, V2ConnectionTo, + Agent, ByteStreams, Client, Conductor, ConnectTo, DynConnectTo, Error, JsonRpcRequest, + JsonRpcResponse, NullRun, Proxy, Responder, V2ConnectionTo, mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, role, schema::{ProtocolVersion, v2}, @@ -22,6 +22,16 @@ use serde_json::json; use tokio::io::duplex; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, JsonRpcRequest)] +#[request(method = "_test/probe", response = ProbeResponse)] +struct ProbeRequest {} + +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize, JsonRpcResponse)] +struct ProbeResponse { + #[serde(rename = "resultType")] + result_type: String, +} + fn implementation(name: &str) -> v2::Implementation { v2::Implementation::new(name, env!("CARGO_PKG_VERSION")) } @@ -40,7 +50,7 @@ fn existing_server() -> v2::McpServer { #[derive(Debug, PartialEq, Eq)] struct ObservedMcpContext { server_id: String, - connection_id: String, + request_id: String, } struct RecordingMcpConnect { @@ -58,9 +68,9 @@ impl McpServerConnect for RecordingMcpConnect { .server_id() .expect("the global MCP server should be attached through ACP") .to_string(), - connection_id: context - .connection_id() - .expect("an attached MCP connection should have an ID") + request_id: context + .request_id() + .expect("an attached MCP request should have an ID") .to_string(), }); DynConnectTo::new(PendingMcpComponent) @@ -73,6 +83,14 @@ impl ConnectTo for PendingMcpComponent { async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { role::mcp::Server .builder() + .on_receive_request( + async |_request: ProbeRequest, responder: Responder, _connection| { + responder.respond(ProbeResponse { + result_type: "complete".to_owned(), + }) + }, + agent_client_protocol::on_receive_request!(), + ) .connect_with(client, async |_connection| { std::future::pending::>().await }) @@ -214,14 +232,17 @@ impl ConnectTo for RecordingAgent { let mcp_connection = connection.clone(); connection.spawn(async move { let result = async { - let connected = mcp_connection - .send_request(v2::ConnectMcpRequest::new(server_id)) - .block_task() - .await?; mcp_connection - .send_request(v2::DisconnectMcpRequest::new( - connected.connection_id, - )) + .send_request(v2::MessageMcpRequest::new( + server_id, + v2::McpRequestId::new("global-v2-probe"), + "_test/probe", + ).params(serde_json::Map::from_iter([ + ("_meta".to_owned(), json!({ + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + })), + ]))) .block_task() .await?; Ok(()) @@ -335,7 +356,7 @@ async fn v2_global_mcp_attachment_preserves_setup_and_continues_handler_chain() tokio::time::timeout(std::time::Duration::from_secs(2), round_trip_rx.next()) .await - .expect("global MCP connect/disconnect round trip should not hang") + .expect("global MCP request should not hang") .ok_or_else(|| Error::internal_error().data("MCP round-trip channel closed"))??; connection @@ -369,8 +390,8 @@ async fn v2_global_mcp_attachment_preserves_setup_and_continues_handler_chain() assert_eq!(mcp_contexts.len(), 1); assert_eq!(mcp_contexts[0].server_id, server_ids[0].to_string()); assert!( - !mcp_contexts[0].connection_id.is_empty(), - "the global MCP connection should receive a connection ID" + mcp_contexts[0].request_id == "global-v2-probe", + "the global MCP request should retain its logical request ID" ); Ok(()) } diff --git a/src/agent-client-protocol-conductor/tests/request_cancellation.rs b/src/agent-client-protocol-conductor/tests/request_cancellation.rs index 34eac429..a2449cfd 100644 --- a/src/agent-client-protocol-conductor/tests/request_cancellation.rs +++ b/src/agent-client-protocol-conductor/tests/request_cancellation.rs @@ -21,12 +21,12 @@ use std::time::Duration; use agent_client_protocol::DynConnectTo; use agent_client_protocol::schema::ProtocolVersion; use agent_client_protocol::schema::v1::{ - CancelRequestNotification, ConnectMcpRequest, ContentBlock, ContentChunk, InitializeRequest, - InitializeResponse, McpServer as SchemaMcpServer, McpServerAcpId, NewSessionRequest, - NewSessionResponse, PermissionOption, PermissionOptionKind, PromptRequest, PromptResponse, - RequestPermissionOutcome, RequestPermissionRequest, RequestPermissionResponse, - SelectedPermissionOutcome, SessionId, SessionNotification, SessionUpdate, StopReason, - ToolCallUpdate, ToolCallUpdateFields, + CancelRequestNotification, ContentBlock, ContentChunk, InitializeRequest, InitializeResponse, + McpRequestId, McpServer as SchemaMcpServer, McpServerAcpId, MessageMcpNotification, + MessageMcpRequest, MessageMcpResponse, NewSessionRequest, NewSessionResponse, PermissionOption, + PermissionOptionKind, PromptRequest, PromptResponse, RequestId, RequestPermissionOutcome, + RequestPermissionRequest, RequestPermissionResponse, SelectedPermissionOutcome, SessionId, + SessionNotification, SessionUpdate, StopReason, ToolCallUpdate, ToolCallUpdateFields, }; use agent_client_protocol::{ Agent, ByteStreams, Client, Conductor, ConnectTo, ConnectionTo, Error, JsonRpcRequest, @@ -204,12 +204,12 @@ async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Erro .await }); - let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let initialize = cx .send_request(InitializeRequest::new(ProtocolVersion::V1)) .block_task() @@ -220,6 +220,7 @@ async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Erro message: "park".into(), }); let client_request_id = request.id().clone(); + let parked_id = next_with_timeout(&mut parked_id_rx).await; request.cancel()?; // The cancellation reaches the agent hop by hop, and the @@ -242,7 +243,7 @@ async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Erro .await?; assert_eq!(barrier.result, "echo: barrier"); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }, ) .await @@ -250,10 +251,10 @@ async fn client_cancellation_propagates_hop_by_hop_to_agent() -> Result<(), Erro .await .expect("test timed out") .expect("client failed"); + let (client_request_id, parked_id) = client_result; // The agent saw exactly one `$/cancel_request`, for the request ID on // its own connection. - let parked_id = next_with_timeout(&mut parked_id_rx).await; assert_ne!( parked_id, client_request_id, "each hop must re-issue the request under its own ID" @@ -277,6 +278,9 @@ async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Erro let (client_cancel_tx, mut client_cancel_rx) = mpsc::unbounded(); // The JSON-RPC id of the parked request, as seen by the client. let (parked_id_tx, mut parked_id_rx) = mpsc::unbounded(); + let (parked_tx, parked_rx) = tokio::sync::oneshot::channel(); + let parked_tx = Arc::new(Mutex::new(Some(parked_tx))); + let parked_rx = Arc::new(Mutex::new(Some(parked_rx))); let agent = Agent .builder() @@ -287,11 +291,12 @@ async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Erro agent_client_protocol::on_receive_request!(), ) .on_receive_request( - async |request: SimpleRequest, - responder: Responder, - cx: ConnectionTo| { + async move |request: SimpleRequest, + responder: Responder, + cx: ConnectionTo| { if request.message == "trigger reverse cancel" { let connection = cx.clone(); + let parked_rx = parked_rx.lock().unwrap().take().expect("one trigger"); cx.spawn(async move { // Send a request to the client, cancel it, and report // how it concluded as the response to the trigger. @@ -299,6 +304,10 @@ async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Erro connection.send_request(SimpleRequest { message: "park".into(), }); + tokio::time::timeout(Duration::from_secs(10), parked_rx) + .await + .expect("timed out waiting for client to park request") + .expect("client closed parked request channel"); upstream.cancel()?; let error = upstream .block_task() @@ -342,6 +351,13 @@ async fn agent_cancellation_propagates_hop_by_hop_to_client() -> Result<(), Erro cx: ConnectionTo| { assert_eq!(request.message, "park"); parked_id_tx.unbounded_send(responder.id().clone()).unwrap(); + parked_tx + .lock() + .unwrap() + .take() + .expect("one parked request") + .send(()) + .expect("agent still waiting to cancel"); let cancellation = responder.cancellation(); cx.spawn(async move { let response = cancellation @@ -534,7 +550,7 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E .await }); - let client_prompt_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .on_receive_request( @@ -579,7 +595,7 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E ) .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let initialize = cx .send_request(InitializeRequest::new(ProtocolVersion::V1)) .block_task() @@ -598,6 +614,8 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E vec!["park".into()], )); let client_prompt_id = prompt.id().clone(); + let prompt_id = next_with_timeout(&mut prompt_id_rx).await; + let permission_id = next_with_timeout(&mut permission_id_rx).await; prompt.cancel()?; let error = prompt @@ -619,7 +637,7 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E .await?; assert_eq!(barrier.stop_reason, StopReason::EndTurn); - Ok(client_prompt_id) + Ok((client_prompt_id, prompt_id, permission_id)) }, ) .await @@ -627,10 +645,10 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E .await .expect("test timed out") .expect("client failed"); + let (client_prompt_id, prompt_id, permission_id) = client_result; // The agent saw exactly one `$/cancel_request` (for the prompt), with the // ID of the prompt on the conductor-to-agent connection. - let prompt_id = next_with_timeout(&mut prompt_id_rx).await; assert_ne!( prompt_id, client_prompt_id, "each hop must re-issue the request under its own ID" @@ -641,7 +659,6 @@ async fn prompt_cancellation_cascades_through_real_proxy_chain() -> Result<(), E // The client saw exactly one `$/cancel_request` (for the permission // request), with the ID of that request on the client's own connection. - let permission_id = next_with_timeout(&mut permission_id_rx).await; let observed = next_with_timeout(&mut client_cancel_rx).await; assert_eq!(observed, permission_id); assert_no_event(&mut client_cancel_rx); @@ -717,12 +734,12 @@ async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error .await }); - let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let initialize = cx .send_request(InitializeRequest::new(ProtocolVersion::V1)) .block_task() @@ -732,6 +749,7 @@ async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error let request: SentRequest = cx.send_request(NewSessionRequest::new("/park-session")); let client_request_id = request.id().clone(); + let parked_id = next_with_timeout(&mut parked_id_rx).await; request.cancel()?; let error = request @@ -750,7 +768,7 @@ async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error .await?; assert_eq!(session.session_id, SessionId::new("normal-session")); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }, ) .await @@ -758,10 +776,10 @@ async fn session_new_cancellation_propagates_through_proxy() -> Result<(), Error .await .expect("test timed out") .expect("client failed"); + let (client_request_id, parked_id) = client_result; // The agent saw exactly one `$/cancel_request`, for the `session/new` ID // on its own connection. - let parked_id = next_with_timeout(&mut parked_id_rx).await; assert_ne!( parked_id, client_request_id, "each hop must re-issue the request under its own ID" @@ -845,12 +863,12 @@ async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), E .await }); - let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let initialize = cx .send_request(InitializeRequest::new(ProtocolVersion::V1)) .block_task() @@ -860,6 +878,7 @@ async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), E let request: SentRequest = cx.send_request(NewSessionRequest::new("/park-session")); let client_request_id = request.id().clone(); + let parked_id = next_with_timeout(&mut parked_id_rx).await; request.cancel()?; let error = request @@ -876,7 +895,7 @@ async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), E .await?; assert_eq!(session.session_id, SessionId::new("normal-session")); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }, ) .await @@ -884,8 +903,8 @@ async fn proxy_session_helper_cancellation_propagates_to_agent() -> Result<(), E .await .expect("test timed out") .expect("client failed"); + let (client_request_id, parked_id) = client_result; - let parked_id = next_with_timeout(&mut parked_id_rx).await; assert_ne!( parked_id, client_request_id, "each hop must re-issue the request under its own ID" @@ -909,90 +928,101 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - let (probe_barrier_tx, mut probe_barrier_rx) = mpsc::unbounded(); let cancelled_mcp_server_id = Arc::new(Mutex::new(None::)); - let agent = Agent - .builder() - .on_receive_request( - async |initialize: InitializeRequest, responder, _cx: ConnectionTo| { - responder.respond(InitializeResponse::new(initialize.protocol_version)) - }, - agent_client_protocol::on_receive_request!(), - ) - .on_receive_request( - { - let cancelled_mcp_server_id = cancelled_mcp_server_id.clone(); - let parked_id_tx = parked_id_tx.clone(); - let probe_barrier_tx = probe_barrier_tx.clone(); - async move |request: NewSessionRequest, - responder: Responder, - cx: ConnectionTo| { + let agent = + Agent + .builder() + .on_receive_request( + async |initialize: InitializeRequest, responder, _cx: ConnectionTo| { + responder.respond(InitializeResponse::new(initialize.protocol_version)) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + { let cancelled_mcp_server_id = cancelled_mcp_server_id.clone(); let parked_id_tx = parked_id_tx.clone(); let probe_barrier_tx = probe_barrier_tx.clone(); - let advertised_mcp_server_id = advertised_mcp_server_id(&request); - - if request.cwd.ends_with("park-session") { - *cancelled_mcp_server_id + async move |request: NewSessionRequest, + responder: Responder, + cx: ConnectionTo| { + let cancelled_mcp_server_id = cancelled_mcp_server_id.clone(); + let parked_id_tx = parked_id_tx.clone(); + let probe_barrier_tx = probe_barrier_tx.clone(); + let advertised_mcp_server_id = advertised_mcp_server_id(&request); + + if request.cwd.ends_with("park-session") { + *cancelled_mcp_server_id + .lock() + .expect("cancelled MCP ID mutex poisoned") = + Some(advertised_mcp_server_id); + parked_id_tx.unbounded_send(responder.id().clone()).unwrap(); + let cancellation = responder.cancellation(); + cx.spawn(async move { + let response = cancellation + .run_until_cancelled(std::future::pending::< + Result, + >()) + .await; + responder.respond_with_result(response) + })?; + return Ok(()); + } + + responder + .respond(NewSessionResponse::new(SessionId::new("normal-session")))?; + + let stale_server_id = cancelled_mcp_server_id .lock() - .expect("cancelled MCP ID mutex poisoned") = - Some(advertised_mcp_server_id); - parked_id_tx.unbounded_send(responder.id().clone()).unwrap(); - let cancellation = responder.cancellation(); + .expect("cancelled MCP ID mutex poisoned") + .clone() + .expect("cancelled session should have advertised an MCP server"); + let connection = cx.clone(); cx.spawn(async move { - let response = cancellation - .run_until_cancelled(std::future::pending::< - Result, - >()) - .await; - responder.respond_with_result(response) - })?; - return Ok(()); - } - - responder.respond(NewSessionResponse::new(SessionId::new("normal-session")))?; - - let stale_server_id = cancelled_mcp_server_id - .lock() - .expect("cancelled MCP ID mutex poisoned") - .clone() - .expect("cancelled session should have advertised an MCP server"); - let connection = cx.clone(); - cx.spawn(async move { - connection - .send_request(ConnectMcpRequest::new(stale_server_id)) + connection + .send_request(MessageMcpRequest::new( + stale_server_id, + McpRequestId::new("stale-server-probe"), + "ping", + ).params(serde_json::Map::from_iter([ + ("_meta".to_owned(), serde_json::json!({ + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + })), + ]))) .on_receiving_result(async |_| Ok(()))?; - let barrier = connection - .send_request(RequestPermissionRequest::new( - SessionId::new("normal-session"), - ToolCallUpdate::new( - "stale-mcp-probe-barrier", - ToolCallUpdateFields::default(), - ), - vec![PermissionOption::new( - "allow", - "Allow", - PermissionOptionKind::AllowOnce, - )], - )) - .block_task() - .await - .map(|_| ()) - .map_err(|error| i32::from(error.code)); - - probe_barrier_tx.unbounded_send(barrier).unwrap(); - Ok(()) - }) - } - }, - agent_client_protocol::on_receive_request!(), - ) - .on_receive_notification( - async move |cancel: CancelRequestNotification, _cx: ConnectionTo| { - agent_cancel_tx.unbounded_send(cancel.request_id).unwrap(); - Ok(()) - }, - agent_client_protocol::on_receive_notification!(), - ); + let barrier = connection + .send_request(RequestPermissionRequest::new( + SessionId::new("normal-session"), + ToolCallUpdate::new( + "stale-mcp-probe-barrier", + ToolCallUpdateFields::default(), + ), + vec![PermissionOption::new( + "allow", + "Allow", + PermissionOptionKind::AllowOnce, + )], + )) + .block_task() + .await + .map(|_| ()) + .map_err(|error| i32::from(error.code)); + + probe_barrier_tx.unbounded_send(barrier).unwrap(); + Ok(()) + }) + } + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_notification( + async move |cancel: CancelRequestNotification, _cx: ConnectionTo| { + agent_cancel_tx.unbounded_send(cancel.request_id).unwrap(); + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ); let proxy = Proxy.builder().on_receive_request_from( Client, @@ -1053,6 +1083,7 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - let request: SentRequest = cx.send_request(NewSessionRequest::new("/park-session")); let client_request_id = request.id().clone(); + let parked_id = next_with_timeout(&mut parked_id_rx).await; request.cancel()?; let error = request @@ -1070,7 +1101,7 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - assert_eq!(session.session_id, SessionId::new("normal-session")); let probe_barrier = next_with_timeout(&mut probe_barrier_rx).await; - Ok((client_request_id, probe_barrier)) + Ok((client_request_id, parked_id, probe_barrier)) }, ) .await @@ -1078,9 +1109,8 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - .await .expect("test timed out") .expect("client failed"); - let (client_request_id, probe_barrier) = client_result; + let (client_request_id, parked_id, probe_barrier) = client_result; - let parked_id = next_with_timeout(&mut parked_id_rx).await; assert_ne!( parked_id, client_request_id, "each hop must re-issue the request under its own ID" @@ -1100,6 +1130,237 @@ async fn proxy_session_helper_cleans_up_mcp_handlers_after_cancelled_session() - Ok(()) } +#[derive(Clone)] +struct ParkedMcpServer { + started_tx: mpsc::UnboundedSender, + stopped_tx: mpsc::UnboundedSender, + dropped_tx: mpsc::UnboundedSender<()>, + late_tx: mpsc::UnboundedSender>, +} + +impl McpServerConnect for ParkedMcpServer { + fn name(&self) -> String { + "parked-mcp".into() + } + + fn connect(&self, cx: McpConnectionTo) -> DynConnectTo { + assert_eq!( + cx.request_id().map(ToString::to_string).as_deref(), + Some("logical-mcp-request") + ); + DynConnectTo::new(ParkedMcpComponent(self.clone())) + } +} + +struct ParkedMcpComponent(ParkedMcpServer); + +struct ProbeOnDrop { + sender: mpsc::UnboundedSender, + value: Option, +} + +impl Drop for ProbeOnDrop { + fn drop(&mut self) { + if let Some(value) = self.value.take() { + drop(self.sender.unbounded_send(value)); + } + } +} + +impl ConnectTo for ParkedMcpComponent { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + let started_tx = self.0.started_tx; + let stopped_tx = self.0.stopped_tx; + let late_tx = self.0.late_tx; + let _backend_dropped = ProbeOnDrop { + sender: self.0.dropped_tx, + value: Some(()), + }; + role::mcp::Server + .builder() + .on_receive_request( + async move |_request: McpParkRequest, + responder: Responder, + cx: ConnectionTo| { + let id = responder.id().clone(); + let stopped = ProbeOnDrop { + sender: stopped_tx.clone(), + value: Some(id.clone()), + }; + late_tx.unbounded_send(cx.clone()).unwrap(); + started_tx.unbounded_send(id).unwrap(); + let cancellation = responder.cancellation(); + cx.spawn(async move { + // Request-scoped cancellation drops this whole backend, + // rather than sending a second, inner cancellation RPC. + let _stopped = stopped; + let result = cancellation + .run_until_cancelled(std::future::pending::< + Result, + >()) + .await; + responder.respond_with_result(result) + }) + }, + agent_client_protocol::on_receive_request!(), + ) + .connect_to(client) + .await + } +} + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)] +#[request(method = "_test/park", response = McpParkResponse)] +struct McpParkRequest {} + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)] +struct McpParkResponse {} + +#[derive(Debug, Clone, Serialize, Deserialize, agent_client_protocol::JsonRpcNotification)] +#[notification(method = "_test/late")] +struct LateMcpNotification {} + +/// An MCP operation retains its logical ID while each ACP transport hop +/// rewrites the outer JSON-RPC ID. Cancelling the ACP request tears down the +/// per-operation server and must not deliver a late MCP notification. +#[tokio::test] +async fn mcp_request_cancellation_crosses_proxy_and_tears_down_backend() -> Result<(), Error> { + let (started_tx, mut started_rx) = mpsc::unbounded(); + let (stopped_tx, mut stopped_rx) = mpsc::unbounded(); + let (dropped_tx, mut dropped_rx) = mpsc::unbounded(); + let (late_tx, mut late_rx) = mpsc::unbounded(); + let (request_id_tx, mut request_id_rx) = mpsc::unbounded(); + let (result_tx, mut result_rx) = mpsc::unbounded(); + let (notification_tx, mut notification_rx) = mpsc::unbounded(); + let (cancel_gate_tx, cancel_gate_rx) = tokio::sync::oneshot::channel::<()>(); + let cancel_gate = Arc::new(Mutex::new(Some(cancel_gate_rx))); + + let agent = Agent + .builder() + .on_receive_request( + async |request: InitializeRequest, responder, _cx: ConnectionTo| { + responder.respond(InitializeResponse::new(request.protocol_version)) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: NewSessionRequest, + responder: Responder, + cx: ConnectionTo| { + let server_id = advertised_mcp_server_id(&request); + responder.respond(NewSessionResponse::new(SessionId::new( + "mcp-cancel-session", + )))?; + let gate = cancel_gate + .lock() + .unwrap() + .take() + .expect("one MCP operation"); + let connection = cx.clone(); + let request_id_tx = request_id_tx.clone(); + let result_tx = result_tx.clone(); + cx.spawn(async move { + let request = connection.send_request( + MessageMcpRequest::new( + server_id, + McpRequestId::new("logical-mcp-request"), + "_test/park", + ) + .params(serde_json::Map::from_iter([( + "_meta".into(), + serde_json::json!({ + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + }), + )])), + ); + request_id_tx.unbounded_send(request.id().clone()).unwrap(); + gate.await.map_err(Error::into_internal_error)?; + request.cancel()?; + let result: Result = request.block_task().await; + result_tx + .unbounded_send(result.map(|_| ()).map_err(|error| i32::from(error.code))) + .unwrap(); + Ok(()) + }) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_notification( + async move |notification: MessageMcpNotification, _cx: ConnectionTo| { + notification_tx.unbounded_send(notification).unwrap(); + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ); + let proxy = Proxy.builder().with_mcp_server(McpServer::new( + ParkedMcpServer { + started_tx, + stopped_tx, + dropped_tx, + late_tx, + }, + NullRun, + )); + let (editor_write, conductor_read) = duplex(8192); + let (conductor_write, editor_read) = duplex(8192); + let conductor_handle = tokio::spawn(async move { + ConductorImpl::new_agent( + "mcp-cancel-conductor".to_string(), + ProxiesAndAgent::new(agent).proxy(proxy), + ) + .run(ByteStreams::new( + conductor_write.compat_write(), + conductor_read.compat(), + )) + .await + }); + + tokio::time::timeout(Duration::from_secs(30), async move { + Client + .builder() + .connect_with( + ByteStreams::new(editor_write.compat_write(), editor_read.compat()), + async |cx| { + cx.send_request(InitializeRequest::new(ProtocolVersion::V1)) + .block_task() + .await?; + cx.send_request(NewSessionRequest::new( + std::env::current_dir().map_err(Error::into_internal_error)?, + )) + .block_task() + .await?; + let outer_id = next_with_timeout(&mut request_id_rx).await; + let backend_id = next_with_timeout(&mut started_rx).await; + assert_ne!(outer_id, backend_id, "JSON-RPC IDs must be hop-local"); + assert_eq!( + backend_id, + RequestId::Str("logical-mcp-request".to_owned()), + "the inner MCP ID must survive the proxy unchanged" + ); + cancel_gate_tx + .send(()) + .expect("agent still waiting to cancel"); + assert_eq!(next_with_timeout(&mut result_rx).await, Err(-32800)); + assert_eq!(next_with_timeout(&mut stopped_rx).await, backend_id); + next_with_timeout(&mut dropped_rx).await; + let late = next_with_timeout(&mut late_rx).await; + assert!( + late.send_notification(LateMcpNotification {}).is_err(), + "a stopped backend must reject an attempted late notification" + ); + assert_no_event(&mut notification_rx); + Ok(()) + }, + ) + .await + }) + .await + .expect("MCP cancellation timed out")?; + conductor_handle.abort(); + Ok(()) +} + /// `initialize` is rewritten to `_proxy/initialize` at the conductor-to-proxy /// hop and forwarded with a result hook — cancellation must still propagate /// hop by hop, exactly like every other request. @@ -1162,15 +1423,16 @@ async fn initialize_cancellation_propagates_through_proxy() -> Result<(), Error> .await }); - let client_request_id = tokio::time::timeout(Duration::from_secs(30), async move { + let client_result = tokio::time::timeout(Duration::from_secs(30), async move { Client .builder() .connect_with( ByteStreams::new(editor_write.compat_write(), editor_read.compat()), - async |cx| { + async move |cx| { let request: SentRequest = cx.send_request(InitializeRequest::new(ProtocolVersion::V1)); let client_request_id = request.id().clone(); + let parked_id = next_with_timeout(&mut parked_id_rx).await; request.cancel()?; let error = request @@ -1187,7 +1449,7 @@ async fn initialize_cancellation_propagates_through_proxy() -> Result<(), Error> .await?; assert_eq!(initialize.protocol_version, ProtocolVersion::V1); - Ok(client_request_id) + Ok((client_request_id, parked_id)) }, ) .await @@ -1195,10 +1457,10 @@ async fn initialize_cancellation_propagates_through_proxy() -> Result<(), Error> .await .expect("test timed out") .expect("client failed"); + let (client_request_id, parked_id) = client_result; // The agent saw exactly one `$/cancel_request`, for the `initialize` ID // on its own connection. - let parked_id = next_with_timeout(&mut parked_id_rx).await; assert_ne!( parked_id, client_request_id, "each hop must re-issue the request under its own ID" diff --git a/src/agent-client-protocol-conductor/tests/scoped_mcp_server.rs b/src/agent-client-protocol-conductor/tests/scoped_mcp_server.rs index d0e3a645..5e3eab5c 100644 --- a/src/agent-client-protocol-conductor/tests/scoped_mcp_server.rs +++ b/src/agent-client-protocol-conductor/tests/scoped_mcp_server.rs @@ -39,7 +39,7 @@ async fn test_scoped_mcp_server_through_proxy() -> Result<(), agent_client_proto .await?; expect_test::expect![[r#" - "OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"2\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" + "OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"2\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" "#]].assert_debug_eq(&result); Ok(()) @@ -84,7 +84,7 @@ async fn test_scoped_mcp_server_through_session() -> Result<(), agent_client_pro .await?; expect_test::expect![[r#" - "OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"2\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" + "OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"2\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" "#]].assert_debug_eq(&result); Ok(()) diff --git a/src/agent-client-protocol-conductor/tests/standalone_mcp_server.rs b/src/agent-client-protocol-conductor/tests/standalone_mcp_server.rs index 8de40f55..33118cb1 100644 --- a/src/agent-client-protocol-conductor/tests/standalone_mcp_server.rs +++ b/src/agent-client-protocol-conductor/tests/standalone_mcp_server.rs @@ -41,7 +41,7 @@ fn create_test_server() -> McpServer HTTP polyfill -> ACP/conductor -> rmcp server. + +use std::{ + future::Future, + path::PathBuf, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, Client, Error, + mcp_server::McpServer, + schema::{ + ProtocolVersion, + v1::{ + AgentCapabilities, InitializeRequest, InitializeResponse, McpCapabilities, + McpServer as AcpMcpServer, NewSessionRequest, NewSessionResponse, SessionCapabilities, + }, + }, +}; +use agent_client_protocol_conductor::{ConductorImpl, ProxiesAndAgent}; +use agent_client_protocol_polyfill::mcp_over_acp::McpOverAcpPolyfill; +use agent_client_protocol_rmcp::McpServerExt as _; +use rmcp::{ + ErrorData, RoleServer, ServerHandler, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, ClientCapabilities, ClientConfig, + Implementation, InputRequiredResult, ProtocolVersion as McpVersion, ServerCapabilities, + ServerConfig, SubscriptionFilter, Tool, ToolAnnotations, + }, + service::{ClientLifecycleMode, ClientServiceExt, RequestContext, SubscriptionContext}, + transport::{ + StreamableHttpClientTransport, streamable_http_client::StreamableHttpClientTransportConfig, + }, +}; +use serde_json::{Value, json}; +use tokio::sync::mpsc; + +const TIMEOUT: Duration = Duration::from_secs(15); +const STATE: &str = "opaque/http/retry?keep=exact"; + +struct RealService { + listening: mpsc::UnboundedSender<()>, + stopped: mpsc::UnboundedSender<()>, + lists: Arc, +} + +struct NotifyStopped(mpsc::UnboundedSender<()>); + +impl Drop for NotifyStopped { + fn drop(&mut self) { + let _ = self.0.send(()); + } +} + +impl ServerHandler for RealService { + fn get_info(&self) -> ServerConfig { + ServerConfig::new( + ServerCapabilities::builder() + .enable_tools() + .enable_tool_list_changed() + .build(), + ) + } + + fn list_tools( + &self, + _request: Option, + _cx: RequestContext, + ) -> impl Future> + Send { + self.lists.fetch_add(1, Ordering::SeqCst); + let schema = json!({"type": "object"}).as_object().unwrap().clone(); + let annotated_schema = json!({ + "type": "object", + "properties": {"region": {"type": "string", "x-mcp-header": "Region"}} + }) + .as_object() + .unwrap() + .clone(); + std::future::ready(Ok(rmcp::model::ListToolsResult::with_all_items(vec![ + Tool::new("retry", "MRTR round trip", schema.clone()), + Tool::new("annotated", "Direct call", annotated_schema).with_annotations( + ToolAnnotations::from_raw(Some("Annotated".into()), Some(true), None, None, None), + ), + ]))) + } + + fn call_tool( + &self, + request: CallToolRequestParams, + cx: RequestContext, + ) -> impl Future> + Send { + std::future::ready( + match (request.name.as_ref(), request.request_state.as_deref()) { + ("retry", None) => { + let inputs = serde_json::from_value(json!({"confirmation": { + "method": "elicitation/create", + "params": {"mode": "form", "message": "Confirm", + "requestedSchema": {"type": "object", + "properties": {"approved": {"type": "boolean"}}}} + }})) + .expect("valid input request"); + Ok(InputRequiredResult::new(Some(inputs), Some(STATE.into())).into()) + } + ("retry", Some(STATE)) => Ok(CallToolResult::structured(json!({ + "state": request.request_state.clone(), + "responses": request.input_responses, + "marker": cx.meta.get("example/marker"), + })) + .into()), + ("annotated", None) => { + Ok(CallToolResult::structured(json!({"direct": true})).into()) + } + _ => Err(ErrorData::invalid_params( + "unknown tool or retry state", + None, + )), + }, + ) + } + + fn accepted_subscription_filter( + &self, + requested: &SubscriptionFilter, + ) -> Option { + Some(requested.clone()) + } + + async fn listen(&self, cx: SubscriptionContext) -> Result<(), ErrorData> { + // The owned adapter may cancel by dropping this future before it polls + // cx.cancelled() again. Observe actual cleanup, not a cooperative branch. + let _stopped = NotifyStopped(self.stopped.clone()); + cx.sink() + .notify_tool_list_changed() + .await + .map_err(|e| ErrorData::internal_error(e.to_string(), None))?; + let _ = self.listening.send(()); + cx.cancelled().await; + Ok(()) + } +} + +fn marked_call(name: &str, marker: &str) -> CallToolRequestParams { + let mut params = CallToolRequestParams::new(name.to_owned()); + params.meta = Some( + serde_json::from_value(json!({ + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {"elicitation": {"form": {}}}, + "io.modelcontextprotocol/clientInfo": {"name": "http-integration", "version": "1"}, + "example/marker": marker, + })) + .expect("valid request metadata"), + ); + params +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 4)] +async fn real_rmcp_stateless_http_survives_subscription_cancellation() -> Result<(), Error> { + tokio::time::timeout(TIMEOUT, async { + let (endpoint_tx, mut endpoint_rx) = mpsc::unbounded_channel(); + let (listening_tx, mut listening_rx) = mpsc::unbounded_channel(); + let (stopped_tx, mut stopped_rx) = mpsc::unbounded_channel(); + let lists = Arc::new(AtomicUsize::new(0)); + let agent = Agent + .builder() + .on_receive_request( + async |request: InitializeRequest, responder, _cx| { + responder.respond( + InitializeResponse::new(request.protocol_version).agent_capabilities( + AgentCapabilities::new() + .session_capabilities(SessionCapabilities::new()) + .mcp_capabilities(McpCapabilities::new().http(true)), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: NewSessionRequest, responder, _cx| { + let [AcpMcpServer::Http(server)] = request.mcp_servers.as_slice() else { + panic!("expected a single HTTP MCP server declaration") + }; + assert_eq!(server.name, "real-rmcp"); + assert_eq!(server.headers.len(), 1); + assert_eq!(server.headers[0].name, "Authorization"); + endpoint_tx + .send((server.url.clone(), server.headers[0].value.clone())) + .expect("client still waiting for HTTP declaration"); + responder.respond(NewSessionResponse::new("real-http-session")) + }, + agent_client_protocol::on_receive_request!(), + ); + + Client.builder().connect_with( + ConductorImpl::new_agent( + "http-bridge", + ProxiesAndAgent::new(agent).proxy(McpOverAcpPolyfill::http()), + ), + async move |cx| { + cx.send_request(InitializeRequest::new(ProtocolVersion::V1)) + .block_task().await?; + let service = Arc::new(RealService { + listening: listening_tx, + stopped: stopped_tx, + lists: lists.clone(), + }); + cx.build_session(PathBuf::from("/tmp")) + .with_mcp_server(McpServer::::from_rmcp( + "real-rmcp", move || service.clone(), + ))? + .block_task() + .run_until(async move |_session| { + let (url, bearer) = endpoint_rx.recv().await.expect("HTTP declaration"); + let headers = [( + "Authorization".parse().expect("header name"), + bearer.parse().expect("header value"), + )].into_iter().collect(); + let transport = StreamableHttpClientTransport::from_config( + StreamableHttpClientTransportConfig::with_uri(url) + .custom_headers(headers), + ); + let config = ClientConfig::new( + serde_json::from_value::( + json!({"elicitation": {"form": {}}}), + ).expect("valid capabilities"), + Implementation::new("http-integration", "1"), + ).with_protocol_version(McpVersion::V_2026_07_28); + let client = config.serve_with_lifecycle( + transport, + ClientLifecycleMode::Discover { + preferred_versions: vec![McpVersion::V_2026_07_28], + }, + ).await.map_err(Error::into_internal_error)?; + + // A direct call before any list also tests absence of hidden lists. + let first = client.call_tool_once(marked_call("retry", "first")) + .await.map_err(Error::into_internal_error)?; + let CallToolResponse::InputRequired(first) = first else { + panic!("expected input_required, got {first:?}"); + }; + assert_eq!(lists.load(Ordering::SeqCst), 0, "no hidden tools/list"); + assert_eq!(first.request_state.as_deref(), Some(STATE)); + assert_eq!( + serde_json::to_value(&first.input_requests) + .map_err(Error::into_internal_error)?["confirmation"]["method"], + "elicitation/create" + ); + let responses: Value = + json!({"confirmation": {"action": "accept", "content": {"approved": true}}}); + let second = client.call_tool_once( + marked_call("retry", "second") + .with_request_state(first.request_state.expect("opaque state")) + .with_input_responses(serde_json::from_value(responses.clone()) + .map_err(Error::into_internal_error)?), + ).await.map_err(Error::into_internal_error)?; + let CallToolResponse::Complete(second) = second else { + panic!("expected completed retry, got {second:?}"); + }; + assert_eq!(second.structured_content.as_ref().unwrap()["state"], STATE); + assert_eq!(second.structured_content.as_ref().unwrap()["responses"], responses); + assert_eq!(second.structured_content.as_ref().unwrap()["marker"], "second"); + + let filter = SubscriptionFilter::builder().tools_list_changed().build(); + let mut subscription = client.listen(filter.clone()).await + .map_err(Error::into_internal_error)?; + listening_rx.recv().await.expect("subscription service started"); + assert_eq!(subscription.acknowledged(), &filter); + let notification = subscription.next().await + .map_err(Error::into_internal_error)? + .expect("filtered notification"); + let notification_json = serde_json::to_value(¬ification) + .map_err(Error::into_internal_error)?; + assert_eq!(notification_json["method"], "notifications/tools/list_changed"); + assert_eq!( + notification_json["params"]["_meta"]["io.modelcontextprotocol/subscriptionId"], + serde_json::to_value(subscription.id()) + .map_err(Error::into_internal_error)? + ); + + // An active HTTP SSE listen must not block an ordinary POST. + let parallel = client.call_tool_once(marked_call("annotated", "parallel")) + .await.map_err(Error::into_internal_error)?; + assert!(matches!(parallel, CallToolResponse::Complete(_))); + subscription.cancel().await.map_err(Error::into_internal_error)?; + stopped_rx.recv().await.expect("subscription cancelled upstream"); + drop(subscription); + let direct = client.call_tool_once(marked_call("annotated", "after-cancel")) + .await.map_err(Error::into_internal_error)?; + assert!(matches!(direct, CallToolResponse::Complete(_))); + let tools = client.list_tools(None).await.map_err(Error::into_internal_error)?; + let annotated = tools.tools.iter().find(|tool| tool.name == "annotated") + .expect("annotated tool still listed"); + assert_eq!(annotated.annotations.as_ref().unwrap().read_only_hint, Some(true)); + assert_eq!( + annotated.input_schema.get("properties").unwrap()["region"], + json!({"type": "string"}) + ); + assert_eq!(lists.load(Ordering::SeqCst), 1); + client.cancel().await.map_err(Error::into_internal_error)?; + Ok(()) + }) + .await + }, + ).await + }) + .await + .expect("rmcp/HTTP/ACP integration timed out") +} diff --git a/src/agent-client-protocol-conductor/tests/test_mcp_connection_context.rs b/src/agent-client-protocol-conductor/tests/test_mcp_connection_context.rs index 11245a97..57024156 100644 --- a/src/agent-client-protocol-conductor/tests/test_mcp_connection_context.rs +++ b/src/agent-client-protocol-conductor/tests/test_mcp_connection_context.rs @@ -1,8 +1,7 @@ //! Integration tests for the context delivered to ACP-attached MCP tools. //! -//! This verifies that an attached tool receives both identifiers defined by the -//! native MCP-over-ACP lifecycle: the server ID advertised during session setup -//! and the connection ID created by `mcp/connect`. +//! This verifies that an attached tool receives the server ID advertised during +//! session setup and the logical MCP request ID assigned to the operation. use agent_client_protocol::RunWithConnectionTo; use agent_client_protocol::mcp_server::McpServer; @@ -20,7 +19,7 @@ struct EchoInput {} #[derive(Debug, Serialize, Deserialize, JsonSchema)] struct EchoOutput { server_id: String, - connection_id: String, + request_id: String, } fn create_echo_proxy() -> DynConnectTo { @@ -28,15 +27,15 @@ fn create_echo_proxy() -> DynConnectTo { .instructions("Test MCP server with a connection-context echo tool") .tool_fn_mut( "echo", - "Returns the current MCP connection context", + "Returns the current MCP request context", async |_input: EchoInput, context| { Ok(EchoOutput { server_id: context .server_id() .expect("tool is attached through ACP") .to_string(), - connection_id: context - .connection_id() + request_id: context + .request_id() .expect("tool is attached through ACP") .to_string(), }) @@ -88,7 +87,7 @@ async fn test_list_tools_from_mcp_server() -> Result<(), agent_client_protocol:: expect![[r" Available tools: - - echo: Returns the current MCP connection context"]] + - echo: Returns the current MCP request context"]] .assert_eq(&result); Ok(()) @@ -115,14 +114,10 @@ async fn test_acp_identifiers_are_delivered_to_mcp_tools() let server_id = regex::Regex::new(r#""server_id":\s*String\("mcp-server:[0-9a-f-]+"\)"#) .expect("valid server ID regex"); - let connection_id = - regex::Regex::new(r#""connection_id":\s*String\("mcp-over-acp-connection:[0-9a-f-]+"\)"#) - .expect("valid connection ID regex"); + let request_id = regex::Regex::new(r#""request_id":\s*String\("[^"]+"\)"#) + .expect("valid logical request ID regex"); assert!(server_id.is_match(&result), "unexpected result: {result}"); - assert!( - connection_id.is_match(&result), - "unexpected result: {result}" - ); + assert!(request_id.is_match(&result), "unexpected result: {result}"); Ok(()) } diff --git a/src/agent-client-protocol-conductor/tests/test_tool_fn.rs b/src/agent-client-protocol-conductor/tests/test_tool_fn.rs index 55203413..0d601fb9 100644 --- a/src/agent-client-protocol-conductor/tests/test_tool_fn.rs +++ b/src/agent-client-protocol-conductor/tests/test_tool_fn.rs @@ -74,8 +74,150 @@ async fn test_tool_fn_greet() -> Result<(), agent_client_protocol::Error> { .await?; expect_test::expect![[r#" - "OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"\\\"Hello, World!\\\"\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" + "OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"\\\"Hello, World!\\\"\", meta: None, annotations: None })], structured_content: None, is_error: Some(false), meta: None }" "#]].assert_debug_eq(&result); Ok(()) } + +/// A cancelled call must not poison the mutable runner, and queued work whose +/// result receiver has gone away must never enter the user's closure. +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn cancelled_tool_fn_mut_keeps_acp_alive() -> Result<(), agent_client_protocol::Error> { + use agent_client_protocol::{ + Agent, Client, Error, Responder, V2ConnectionTo, + schema::{ProtocolVersion, v2}, + }; + use std::{ + sync::{Arc, Mutex}, + time::Duration, + }; + use tokio::sync::oneshot; + + #[derive(Debug, Deserialize, Serialize, JsonSchema)] + struct Input { + name: String, + } + let (started_tx, started_rx) = oneshot::channel(); + let started = Arc::new(Mutex::new(Some(started_tx))); + let calls = Arc::new(Mutex::new(Vec::::new())); + let (result_tx, result_rx) = oneshot::channel(); + let invocation = Arc::new(Mutex::new(Some((started_rx, result_tx)))); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("runner-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new() + .session(v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + )), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let [v2::McpServer::Acp(declaration)] = request.mcp_servers.as_slice() else { + panic!("expected one ACP MCP server") + }; + let server = declaration.server_id.clone(); + let (started_rx, result_tx) = + invocation.lock().unwrap().take().expect("one session"); + let call_cx = cx.clone(); + cx.spawn(async move { + let result = async { + let make_request = |id: &str, name: &str| { + let params = serde_json::json!({ + "name": "hold", + "arguments": {"name": name}, + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } + }); + v2::MessageMcpRequest::new(server.clone(), id.to_owned(), "tools/call") + .params(params.as_object().expect("object params").clone()) + }; + let running = call_cx.send_request(make_request("running", "running")); + started_rx.await.map_err(Error::into_internal_error)?; + let queued = call_cx.send_request(make_request("queued", "queued")); + tokio::task::yield_now().await; + queued.cancel()?; + running.cancel()?; + for request in [running, queued] { + let failure = + request.block_task().await.expect_err("cancelled request"); + assert_eq!(i32::from(failure.code), -32800); + } + let healthy = call_cx + .send_request(make_request("after", "after")) + .block_task() + .await?; + let v2::MessageMcpResponse::Result { result, .. } = healthy else { + panic!("healthy tool call did not produce an MCP result") + }; + assert_eq!(result["isError"], false, "healthy result: {result}"); + Ok::<_, Error>(()) + } + .await; + let _sent = result_tx.send(result); + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new( + "runner-session", + ))) + }, + agent_client_protocol::on_receive_request!(), + ); + tokio::time::timeout( + Duration::from_secs(10), + Client.v2().connect_with(agent, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("runner-client", "1"), + )) + .block_task() + .await?; + let recorded = calls.clone(); + let started = started.clone(); + let server = McpServer::::builder("runner") + .tool_fn_mut( + "hold", + "Hold a mutable runner", + async move |input: Input, _cx| { + recorded.lock().unwrap().push(input.name.clone()); + if input.name == "running" { + if let Some(tx) = started.lock().unwrap().take() { + let _sent = tx.send(()); + } + std::future::pending::<()>().await; + } + Ok(serde_json::json!({"value": input.name})) + }, + agent_client_protocol::tool_fn_mut!(), + ) + .build(); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)? + .start_session() + .block_task() + .await?; + result_rx.await.map_err(Error::into_internal_error)??; + assert_eq!(*calls.lock().unwrap(), ["running", "after"]); + Ok::<_, Error>(()) + }), + ) + .await + .expect("cancelled MCP tool call timed out") +} diff --git a/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs b/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs index 3987983a..1582ba90 100644 --- a/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs +++ b/src/agent-client-protocol-conductor/tests/trace_client_mcp_server.rs @@ -33,7 +33,8 @@ use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; /// - Replaces UUIDs with sequential IDs (id:0, id:1, etc.) /// - Replaces session IDs with "session:0", etc. /// - Replaces loopback HTTP endpoints with "http:endpoint:0", etc. -/// - Replaces MCP server and connection IDs with stable sequential IDs +/// - Replaces MCP server and logical request IDs with stable sequential IDs +/// - Redacts runtime-generated HTTP Authorization headers struct EventNormalizer { id_map: HashMap, next_id: usize, @@ -43,8 +44,8 @@ struct EventNormalizer { next_endpoint: usize, server_map: HashMap, next_server: usize, - connection_map: HashMap, - next_connection: usize, + request_map: HashMap, + next_request: usize, } impl EventNormalizer { @@ -58,8 +59,8 @@ impl EventNormalizer { next_endpoint: 0, server_map: HashMap::new(), next_server: 0, - connection_map: HashMap::new(), - next_connection: 0, + request_map: HashMap::new(), + next_request: 0, } } @@ -116,12 +117,12 @@ impl EventNormalizer { .clone() } - fn normalize_connection_id(&mut self, id: &str) -> String { - self.connection_map + fn normalize_request_id(&mut self, id: &str) -> String { + self.request_map .entry(id.to_string()) .or_insert_with(|| { - let n = format!("connection:{}", self.next_connection); - self.next_connection += 1; + let n = format!("request:{}", self.next_request); + self.next_request += 1; n }) .clone() @@ -131,6 +132,10 @@ impl EventNormalizer { fn normalize_json(&mut self, value: serde_json::Value) -> serde_json::Value { match value { serde_json::Value::Object(map) => { + let is_authorization_header = map + .get("name") + .and_then(|v| v.as_str()) + .is_some_and(|name| name.eq_ignore_ascii_case("authorization")); let normalized: serde_json::Map = map .into_iter() .map(|(k, v)| { @@ -158,12 +163,14 @@ impl EventNormalizer { } else { self.normalize_json(v) } - } else if k == "connectionId" { + } else if k == "requestId" { if let serde_json::Value::String(s) = &v { - serde_json::Value::String(self.normalize_connection_id(s)) + serde_json::Value::String(self.normalize_request_id(s)) } else { self.normalize_json(v) } + } else if is_authorization_header && k == "value" { + serde_json::Value::String("[REDACTED]".into()) } else { self.normalize_json(v) }; @@ -360,6 +367,7 @@ async fn test_trace_client_mcp_server() -> Result<(), agent_client_protocol::Err to: "Client", id: String("id:0"), is_error: false, + error_domain: None, payload: Object { "protocolVersion": Number(1), "agentCapabilities": Object { @@ -423,6 +431,7 @@ async fn test_trace_client_mcp_server() -> Result<(), agent_client_protocol::Err to: "Client", id: String("id:1"), is_error: false, + error_domain: None, payload: Object { "sessionId": String("session:0"), "modes": Object { @@ -500,7 +509,7 @@ async fn test_trace_client_mcp_server() -> Result<(), agent_client_protocol::Err "sessionUpdate": String("agent_message_chunk"), "content": Object { "type": String("text"), - "text": String("OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"{\\\"echoed\\\":\\\"Client echoes: Hello from client test!\\\",\\\"call_number\\\":1}\", meta: None, annotations: None })], structured_content: Some(Object {\"echoed\": String(\"Client echoes: Hello from client test!\"), \"call_number\": Number(1)}), is_error: Some(false), meta: None }"), + "text": String("OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"{\\\"echoed\\\":\\\"Client echoes: Hello from client test!\\\",\\\"call_number\\\":1}\", meta: None, annotations: None })], structured_content: Some(Object {\"echoed\": String(\"Client echoes: Hello from client test!\"), \"call_number\": Number(1)}), is_error: Some(false), meta: None }"), }, "messageId": String("testy-message-end-turn-1"), }, @@ -514,6 +523,7 @@ async fn test_trace_client_mcp_server() -> Result<(), agent_client_protocol::Err to: "Client", id: String("id:2"), is_error: false, + error_domain: None, payload: Object { "stopReason": String("end_turn"), }, diff --git a/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs b/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs index 48366ebf..9bede599 100644 --- a/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs +++ b/src/agent-client-protocol-conductor/tests/trace_mcp_tool_call.rs @@ -32,7 +32,8 @@ use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; /// - Replaces UUIDs with sequential IDs (id:0, id:1, etc.) /// - Replaces session IDs with "session:0", etc. /// - Replaces loopback HTTP endpoints with "http:endpoint:0", etc. -/// - Replaces MCP server and connection IDs with stable sequential IDs +/// - Replaces MCP server and logical request IDs with stable sequential IDs +/// - Redacts runtime-generated HTTP Authorization headers struct EventNormalizer { id_map: HashMap, next_id: usize, @@ -42,8 +43,8 @@ struct EventNormalizer { next_endpoint: usize, server_map: HashMap, next_server: usize, - connection_map: HashMap, - next_connection: usize, + request_map: HashMap, + next_request: usize, } impl EventNormalizer { @@ -57,8 +58,8 @@ impl EventNormalizer { next_endpoint: 0, server_map: HashMap::new(), next_server: 0, - connection_map: HashMap::new(), - next_connection: 0, + request_map: HashMap::new(), + next_request: 0, } } @@ -115,12 +116,12 @@ impl EventNormalizer { .clone() } - fn normalize_connection_id(&mut self, id: &str) -> String { - self.connection_map + fn normalize_request_id(&mut self, id: &str) -> String { + self.request_map .entry(id.to_string()) .or_insert_with(|| { - let n = format!("connection:{}", self.next_connection); - self.next_connection += 1; + let n = format!("request:{}", self.next_request); + self.next_request += 1; n }) .clone() @@ -130,6 +131,10 @@ impl EventNormalizer { fn normalize_json(&mut self, value: serde_json::Value) -> serde_json::Value { match value { serde_json::Value::Object(map) => { + let is_authorization_header = map + .get("name") + .and_then(|v| v.as_str()) + .is_some_and(|name| name.eq_ignore_ascii_case("authorization")); let normalized: serde_json::Map = map .into_iter() .map(|(k, v)| { @@ -157,12 +162,14 @@ impl EventNormalizer { } else { self.normalize_json(v) } - } else if k == "connectionId" { + } else if k == "requestId" { if let serde_json::Value::String(s) = &v { - serde_json::Value::String(self.normalize_connection_id(s)) + serde_json::Value::String(self.normalize_request_id(s)) } else { self.normalize_json(v) } + } else if is_authorization_header && k == "value" { + serde_json::Value::String("[REDACTED]".into()) } else { self.normalize_json(v) }; @@ -317,7 +324,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> // Snapshot the trace events // This should show: // 1. Client -> Agent: initialize, session/new, session/prompt (left-to-right) - // 2. Agent -> MCP Server: tools/call (right-to-left, the key part!) + // 2. Agent -> MCP Server: discovery/list/call with per-request MCP metadata // 3. MCP Server -> Agent: response // 4. Agent -> Client: notification + response expect![[r#" @@ -377,6 +384,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Proxy(0)", id: String("id:1"), is_error: false, + error_domain: None, payload: Object { "protocolVersion": Number(1), "agentCapabilities": Object { @@ -419,6 +427,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Client", id: String("id:0"), is_error: false, + error_domain: None, payload: Object { "protocolVersion": Number(1), "agentCapabilities": Object { @@ -490,32 +499,6 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> }, }, ), - Request( - RequestEvent { - ts: 0.0, - protocol: Acp, - from: "Proxy(1)", - to: "Proxy(0)", - id: String("id:4"), - method: "mcp/connect", - session: None, - params: Object { - "serverId": String("server:0"), - }, - }, - ), - Response( - ResponseEvent { - ts: 0.0, - from: "Proxy(0)", - to: "Proxy(1)", - id: String("id:4"), - is_error: false, - payload: Object { - "connectionId": String("connection:0"), - }, - }, - ), Response( ResponseEvent { ts: 0.0, @@ -523,6 +506,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Proxy(0)", id: String("id:3"), is_error: false, + error_domain: None, payload: Object { "sessionId": String("session:0"), "modes": Object { @@ -573,6 +557,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> to: "Client", id: String("id:2"), is_error: false, + error_domain: None, payload: Object { "sessionId": String("session:0"), "modes": Object { @@ -622,7 +607,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> protocol: Acp, from: "Client", to: "Proxy(0)", - id: String("id:5"), + id: String("id:4"), method: "session/prompt", session: None, params: Object { @@ -642,7 +627,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> protocol: Acp, from: "Proxy(0)", to: "Proxy(1)", - id: String("id:6"), + id: String("id:5"), method: "session/prompt", session: None, params: Object { @@ -662,15 +647,17 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> protocol: Mcp, from: "Proxy(1)", to: "Proxy(0)", - id: String("id:7"), - method: "initialize", + id: String("id:6"), + method: "server/discover", session: None, params: Object { - "protocolVersion": String("2025-11-25"), - "capabilities": Object {}, - "clientInfo": Object { - "name": String("rmcp"), - "version": String("3.4.0"), + "_meta": Object { + "io.modelcontextprotocol/protocolVersion": String("2026-07-28"), + "io.modelcontextprotocol/clientInfo": Object { + "name": String("testy"), + "version": String("0.11.0"), + }, + "io.modelcontextprotocol/clientCapabilities": Object {}, }, }, }, @@ -680,43 +667,46 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> ts: 0.0, from: "Proxy(0)", to: "Proxy(1)", - id: String("id:7"), + id: String("id:6"), is_error: false, + error_domain: None, payload: Object { - "protocolVersion": String("2025-11-25"), + "resultType": String("complete"), + "supportedVersions": Array [ + String("2026-07-28"), + ], "capabilities": Object { "tools": Object {}, }, - "serverInfo": Object { - "name": String("rmcp"), - "version": String("3.4.0"), - }, "instructions": String("A simple test MCP server with an echo tool"), + "ttlMs": Number(0), + "cacheScope": String("private"), + "_meta": Object { + "io.modelcontextprotocol/serverInfo": Object { + "name": String("rmcp"), + "version": String("3.4.0"), + }, + }, }, }, ), - Notification( - NotificationEvent { - ts: 0.0, - protocol: Mcp, - from: "Proxy(1)", - to: "Proxy(0)", - method: "notifications/initialized", - session: None, - params: Null, - }, - ), Request( RequestEvent { ts: 0.0, protocol: Mcp, from: "Proxy(1)", to: "Proxy(0)", - id: String("id:8"), + id: String("id:7"), method: "tools/call", session: None, params: Object { "_meta": Object { + "io.modelcontextprotocol/protocolVersion": String("2026-07-28"), + "io.modelcontextprotocol/clientInfo": Object { + "name": String("testy"), + "version": String("0.11.0"), + }, + "io.modelcontextprotocol/clientCapabilities": Object {}, "progressToken": Number(0), }, "name": String("echo"), @@ -731,9 +721,11 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> ts: 0.0, from: "Proxy(0)", to: "Proxy(1)", - id: String("id:8"), + id: String("id:7"), is_error: false, + error_domain: None, payload: Object { + "resultType": String("complete"), "content": Array [ Object { "type": String("text"), @@ -761,7 +753,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> "sessionUpdate": String("agent_message_chunk"), "content": Object { "type": String("text"), - "text": String("OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"{\\\"result\\\":\\\"Echo: Hello from trace test!\\\"}\", meta: None, annotations: None })], structured_content: Some(Object {\"result\": String(\"Echo: Hello from trace test!\")}), is_error: Some(false), meta: None }"), + "text": String("OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"{\\\"result\\\":\\\"Echo: Hello from trace test!\\\"}\", meta: None, annotations: None })], structured_content: Some(Object {\"result\": String(\"Echo: Hello from trace test!\")}), is_error: Some(false), meta: None }"), }, "messageId": String("testy-message-end-turn-1"), }, @@ -773,8 +765,9 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> ts: 0.0, from: "Proxy(1)", to: "Proxy(0)", - id: String("id:6"), + id: String("id:5"), is_error: false, + error_domain: None, payload: Object { "stopReason": String("end_turn"), }, @@ -794,7 +787,7 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> "sessionUpdate": String("agent_message_chunk"), "content": Object { "type": String("text"), - "text": String("OK: CallToolResult { result_type: None, content: [Text(TextContent { text: \"{\\\"result\\\":\\\"Echo: Hello from trace test!\\\"}\", meta: None, annotations: None })], structured_content: Some(Object {\"result\": String(\"Echo: Hello from trace test!\")}), is_error: Some(false), meta: None }"), + "text": String("OK: CallToolResult { result_type: Some(ResultType(\"complete\")), content: [Text(TextContent { text: \"{\\\"result\\\":\\\"Echo: Hello from trace test!\\\"}\", meta: None, annotations: None })], structured_content: Some(Object {\"result\": String(\"Echo: Hello from trace test!\")}), is_error: Some(false), meta: None }"), }, "messageId": String("testy-message-end-turn-1"), }, @@ -806,8 +799,9 @@ async fn test_trace_mcp_tool_call() -> Result<(), agent_client_protocol::Error> ts: 0.0, from: "Proxy(0)", to: "Client", - id: String("id:5"), + id: String("id:4"), is_error: false, + error_domain: None, payload: Object { "stopReason": String("end_turn"), }, diff --git a/src/agent-client-protocol-conductor/tests/trace_snapshot.rs b/src/agent-client-protocol-conductor/tests/trace_snapshot.rs index 51b48b51..d215ca1b 100644 --- a/src/agent-client-protocol-conductor/tests/trace_snapshot.rs +++ b/src/agent-client-protocol-conductor/tests/trace_snapshot.rs @@ -216,6 +216,7 @@ async fn test_trace_snapshot() -> Result<(), agent_client_protocol::Error> { to: "Client", id: String("id:0"), is_error: false, + error_domain: None, payload: Object { "protocolVersion": Number(1), "agentCapabilities": Object { @@ -273,6 +274,7 @@ async fn test_trace_snapshot() -> Result<(), agent_client_protocol::Error> { to: "Client", id: String("id:1"), is_error: false, + error_domain: None, payload: Object { "sessionId": String("session:0"), "modes": Object { @@ -364,6 +366,7 @@ async fn test_trace_snapshot() -> Result<(), agent_client_protocol::Error> { to: "Client", id: String("id:2"), is_error: false, + error_domain: None, payload: Object { "stopReason": String("end_turn"), }, diff --git a/src/agent-client-protocol-cookbook/src/lib.rs b/src/agent-client-protocol-cookbook/src/lib.rs index 58fba618..e5686a45 100644 --- a/src/agent-client-protocol-cookbook/src/lib.rs +++ b/src/agent-client-protocol-cookbook/src/lib.rs @@ -732,7 +732,7 @@ pub mod global_mcp_server { //! ``` //! //! The `from_rmcp` function takes a factory closure that creates a new server - //! instance. This allows each MCP connection to get a fresh server instance. + //! instance for each MCP request. //! //! # How it works //! @@ -740,13 +740,14 @@ pub mod global_mcp_server { //! handler. It: //! //! 1. Intercepts session setup requests and adds a schema-native - //! `McpServer::Acp` declaration with one connection-scoped server ID. + //! `McpServer::Acp` declaration with one stable server ID. //! V1 injects it into `session/new`, `session/load`, `session/resume`, //! and feature-gated `session/fork`; v2 injects it into //! `session/new`, `session/resume`, and feature-gated `session/fork` //! while preserving unrelated request fields //! 2. Passes the modified request through to the next handler - //! 3. Handles `mcp/connect`, `mcp/message`, and `mcp/disconnect` for that server ID + //! 3. Handles `mcp/message` requests for that server ID. Each operation + //! has its own logical request ID and per-request MCP metadata. //! //! [`McpServer::builder`]: agent_client_protocol_rmcp::McpServerExt::builder //! [`McpServer::from_rmcp`]: agent_client_protocol_rmcp::McpServerExt::from_rmcp diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index 1c899872..3d0730f4 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -4,14 +4,14 @@ use std::{ }; use agent_client_protocol::{ - Agent, Channel, Client, ConnectTo, Error as AcpError, RawJsonRpcMessage, TransportBatchEntry, - TransportFrame, - schema::v1::{RequestId, Response as RpcResponse}, + Agent, BudgetedFrame, Channel, Client, ConnectTo, Error as AcpError, FrameAdmission, + FramePermit, FrameReceiver, FrameSender, RawJsonRpcMessage, RawJsonRpcResponse as RpcResponse, + TransportBatchEntry, TransportFrame, schema::v1::RequestId, }; use async_tungstenite::tungstenite::Message as WsMessage; use futures::{ - Stream, StreamExt, - channel::mpsc::{self, UnboundedSender}, + SinkExt, Stream, StreamExt, + channel::mpsc, future::{BoxFuture, FutureExt}, pin_mut, stream::FuturesUnordered, @@ -20,8 +20,9 @@ use thiserror::Error; use tracing::{debug, error, trace, warn}; use crate::protocol::{ - HEADER_CONNECTION_ID, HEADER_SESSION_ID, is_initialize_request, is_response_only_shape, - method_for_message, method_requires_session_header, session_id_from_message, + HEADER_CONNECTION_ID, HEADER_SESSION_ID, is_cancel_request_message, is_initialize_request, + is_response_only_shape, method_for_message, method_requires_session_header, + session_id_from_message, }; #[derive(Debug, Error)] @@ -123,9 +124,12 @@ impl ConnectTo for HttpClient { } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), AcpError>>) { + fn into_channel_and_future(self) -> (Channel, agent_client_protocol::ConnectionDriver) { let (caller, transport) = Channel::duplex(); - (caller, Box::pin(run(self, transport))) + ( + caller, + agent_client_protocol::ConnectionDriver::new(run(self, transport)), + ) } } @@ -138,15 +142,18 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { rx: mut outgoing, tx: incoming, } = channel; - let (sse_event_tx, mut sse_event_rx) = mpsc::unbounded::(); + let admission = incoming.admission(); + let max_operations = admission.limits().max_queued_frames.max(1); + let (sse_event_tx, mut sse_event_rx) = mpsc::channel::(max_operations); let connection = HttpConnection::new(endpoint, http); let mut state = ClientState { connection: connection.clone(), open_session_streams: HashSet::new(), pending_requests: HashMap::new(), + pending_request_leases: HashMap::new(), incoming, }; - let mut lifecycle = HttpTransportLifecycle::new(connection); + let mut lifecycle = HttpTransportLifecycle::new(connection, admission, max_operations); let mut posts = PostQueues::default(); let mut buffered_outgoing = VecDeque::new(); let mut outgoing_closed = false; @@ -200,13 +207,14 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { let Some(event) = event else { continue; }; - let open_session_ids = state.sessions_to_open_for_responses(&event.frame); - state.deliver_frame(event.frame); + let open_session_ids = state.sessions_to_open_for_responses(event.frame.frame()); + state.deliver_budgeted(event.frame).await?; for session_id in open_session_ids { match lifecycle .start_sse( Some(session_id), sse_event_tx.clone(), + false, SseStartContext { events: &mut sse_event_rx, outgoing: &mut outgoing, @@ -242,7 +250,9 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { } }; - let is_response_only = is_response_only_frame(&frame); + let bypass_ordered = + is_response_only_frame(frame.frame()) || is_cancellation_frame(frame.frame()); + let (frame, permit) = frame.into_parts(); let msg = match frame { TransportFrame::Single(message) => message, frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => { @@ -254,11 +264,26 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { // Response-only batches answer SSE-delivered callbacks and // must not be blocked behind the request they answer. Ok((post, session_ids)) => { + state.attach_pending_permits(&post.pending_requests, &permit); + if let Err(error) = + check_post_capacity(&posts, max_operations, bypass_ordered) + { + break 'transport Err(error); + } + if bypass_ordered { + posts.responses.push_budgeted(post, permit); + } else { + posts.ordered.push_budgeted(post, permit); + } + // The POST registers bounded session mailboxes. Poll it + // while establishing the GET, rather than waiting for + // a GET that cannot succeed before the POST. for session_id in session_ids { match lifecycle .start_sse( Some(session_id), sse_event_tx.clone(), + true, SseStartContext { events: &mut sse_event_rx, outgoing: &mut outgoing, @@ -270,17 +295,17 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { .await { Ok(SseStartOutcome::Established) => {} + Ok(SseStartOutcome::OutgoingClosed) + if buffered_outgoing.is_empty() && posts.is_empty() => + { + break 'transport Ok(()); + } Ok(SseStartOutcome::OutgoingClosed) => { break 'transport Err(sse_setup_blocked_output_error()); } Err(error) => break 'transport Err(error), } } - if is_response_only { - posts.responses.push(post); - } else { - posts.ordered.push(post); - } } Err(error) => { error!("POST failed: {error}"); @@ -302,6 +327,7 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { .start_sse( None, sse_event_tx.clone(), + false, SseStartContext { events: &mut sse_event_rx, outgoing: &mut outgoing, @@ -331,12 +357,34 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { continue; } - if let Some(session_id) = session_id_from_message(&msg) { + let session_id = session_id_from_message(&msg); + if let Err(error) = check_post_capacity(&posts, max_operations, bypass_ordered) { + break Err(error); + } + match state.prepare_post(msg) { + // Responses and cancellation must not be blocked behind a POST + // that may itself be waiting for their delivery. + Ok(post) => { + state.attach_pending_permits(&post.pending_requests, &permit); + if bypass_ordered { + posts.responses.push_budgeted(post, permit); + } else { + posts.ordered.push_budgeted(post, permit); + } + } + Err(e) => { + error!("POST failed: {e}"); + break Err(AcpError::internal_error().data(format!("POST: {e}"))); + } + } + + if let Some(session_id) = session_id { for session_id in state.register_session_streams([session_id]) { match lifecycle .start_sse( Some(session_id), sse_event_tx.clone(), + true, SseStartContext { events: &mut sse_event_rx, outgoing: &mut outgoing, @@ -348,6 +396,11 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { .await { Ok(SseStartOutcome::Established) => {} + Ok(SseStartOutcome::OutgoingClosed) + if buffered_outgoing.is_empty() && posts.is_empty() => + { + break 'transport Ok(()); + } Ok(SseStartOutcome::OutgoingClosed) => { break 'transport Err(sse_setup_blocked_output_error()); } @@ -355,17 +408,6 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { } } } - - match state.prepare_post(msg) { - // Responses answer SSE-delivered callbacks and must not be blocked - // behind a POST that may be waiting for that callback response. - Ok(post) if is_response_only => posts.responses.push(post), - Ok(post) => posts.ordered.push(post), - Err(e) => { - error!("POST failed: {e}"); - break Err(AcpError::internal_error().data(format!("POST: {e}"))); - } - } }; lifecycle.close().await; @@ -383,6 +425,25 @@ fn sse_setup_blocked_output_error() -> AcpError { .data("outgoing channel closed while accepted messages awaited SSE stream establishment") } +fn post_capacity_error() -> AcpError { + AcpError::internal_error().data("HTTP POST operation capacity exceeded") +} + +fn check_post_capacity( + posts: &PostQueues, + max_operations: usize, + bypass_ordered: bool, +) -> Result<(), AcpError> { + // Keep one operation available for callbacks/cancellation even while the + // ordered data POST is waiting for exactly such a response. + let reserved = usize::from(!bypass_ordered && max_operations > 1); + if posts.len() >= max_operations.saturating_sub(reserved).max(1) { + Err(post_capacity_error()) + } else { + Ok(()) + } +} + fn handle_completed_post( state: &mut ClientState, completed: CompletedPost, @@ -403,8 +464,14 @@ fn handle_completed_post( fn queue_response_post( state: &mut ClientState, posts: &mut PostQueues, - frame: TransportFrame, + frame: BudgetedFrame, ) -> Result<(), AcpError> { + check_post_capacity( + posts, + state.incoming.admission().limits().max_queued_frames.max(1), + true, + )?; + let (frame, permit) = frame.into_parts(); let post = match frame { TransportFrame::Single(message) => state.prepare_post(message), frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => { @@ -418,7 +485,8 @@ fn queue_response_post( error!("POST failed: {error}"); AcpError::internal_error().data(format!("POST: {error}")) })?; - posts.responses.push(post); + state.attach_pending_permits(&post.pending_requests, &permit); + posts.responses.push_budgeted(post, permit); Ok(()) } @@ -441,8 +509,27 @@ fn is_response_only_frame(frame: &TransportFrame) -> bool { } } +fn is_cancellation_frame(frame: &TransportFrame) -> bool { + match frame { + TransportFrame::Single(message) => is_cancel_request_message(message), + TransportFrame::Batch(batch) => { + let mut has_cancellation = false; + let only_control = batch.entries().all(|entry| match entry { + TransportBatchEntry::Message(RawJsonRpcMessage::Response(_)) => true, + TransportBatchEntry::Message(message) if is_cancel_request_message(message) => { + has_cancellation = true; + true + } + TransportBatchEntry::Malformed { .. } | TransportBatchEntry::Message(_) => false, + }); + only_control && has_cancellation + } + TransportFrame::Malformed { .. } => false, + } +} + enum HttpLoopEvent { - Outgoing(Option), + Outgoing(Option), SseEvent(Option), SseFailure(SseFailure), Post(CompletedPost), @@ -456,7 +543,7 @@ struct SseFailure { #[derive(Debug)] struct SseMessage { - frame: TransportFrame, + frame: BudgetedFrame, } #[derive(Clone, Debug)] @@ -535,6 +622,9 @@ impl HttpConnection { if let Err(e) = http .delete(endpoint) .header(HEADER_CONNECTION_ID, connection_id) + // A stalled peer must not keep the transport's shutdown (and its + // retained POST/SSE permits) alive indefinitely. + .timeout(std::time::Duration::from_secs(2)) .send() .await { @@ -547,6 +637,8 @@ impl HttpConnection { struct HttpTransportLifecycle { connection: HttpConnection, sse_tasks: SseTasks, + admission: FrameAdmission, + max_tasks: usize, } #[derive(Clone, Copy, Debug, Eq, PartialEq)] @@ -556,25 +648,28 @@ enum SseStartOutcome { } struct SseStartContext<'a> { - events: &'a mut mpsc::UnboundedReceiver, - outgoing: &'a mut mpsc::UnboundedReceiver, - buffered_outgoing: &'a mut VecDeque, + events: &'a mut mpsc::Receiver, + outgoing: &'a mut FrameReceiver, + buffered_outgoing: &'a mut VecDeque, posts: &'a mut PostQueues, state: &'a mut ClientState, } impl HttpTransportLifecycle { - fn new(connection: HttpConnection) -> Self { + fn new(connection: HttpConnection, admission: FrameAdmission, max_tasks: usize) -> Self { Self { connection, sse_tasks: SseTasks::default(), + admission, + max_tasks, } } async fn start_sse( &mut self, session_id: Option, - event_tx: UnboundedSender, + event_tx: mpsc::Sender, + wait_for_post: bool, context: SseStartContext<'_>, ) -> Result { let SseStartContext { @@ -585,15 +680,32 @@ impl HttpTransportLifecycle { state, } = context; let mut establishing = FuturesUnordered::new(); - establishing.push(self.begin_sse(session_id, event_tx.clone())); + let mut session_id = session_id; + if !wait_for_post { + establishing.push(self.begin_sse(session_id.take(), event_tx.clone())?); + } loop { - if establishing.is_empty() { + // A session-scoped POST receives 202 only after its route and + // mailbox are registered. Keep pumping all existing SSE streams, + // callbacks, and POST completions until that admission completes; + // then open the session GET without racing a 409. + if wait_for_post && posts.is_empty() && session_id.is_some() { + establishing.push(self.begin_sse(session_id.take(), event_tx.clone())?); + } + if establishing.is_empty() && session_id.is_none() { return Ok(SseStartOutcome::Established); } let outcome = { let failure = self.sse_tasks.next_failure().fuse(); - let established_next = establishing.next().fuse(); + let established_next = async { + if establishing.is_empty() { + futures::future::pending().await + } else { + establishing.next().await + } + } + .fuse(); let sse_event_next = events.next().fuse(); let outgoing_next = outgoing.next().fuse(); let ordered_post_next = posts.ordered.next_completion().fuse(); @@ -621,24 +733,35 @@ impl HttpTransportLifecycle { return Err(sse_failure_error(self.sse_tasks.next_failure().await)); } SseStartWait::Established(None) => { - return Ok(SseStartOutcome::Established); + if session_id.is_none() { + return Ok(SseStartOutcome::Established); + } } SseStartWait::Failure(failure) => return Err(sse_failure_error(failure)), SseStartWait::SseEvent(Some(event)) => { - let open_session_ids = state.sessions_to_open_for_responses(&event.frame); - state.deliver_frame(event.frame); + let open_session_ids = + state.sessions_to_open_for_responses(event.frame.frame()); + state.deliver_budgeted(event.frame).await?; for session_id in open_session_ids { - establishing.push(self.begin_sse(Some(session_id), event_tx.clone())); + establishing.push(self.begin_sse(Some(session_id), event_tx.clone())?); } } SseStartWait::SseEvent(None) => { return Err(AcpError::internal_error().data("SSE event channel closed")); } SseStartWait::Post(completed) => handle_completed_post(state, completed)?, - SseStartWait::Outgoing(Some(frame)) if is_response_only_frame(&frame) => { + SseStartWait::Outgoing(Some(frame)) + if is_response_only_frame(frame.frame()) + || is_cancellation_frame(frame.frame()) => + { queue_response_post(state, posts, frame)?; } - SseStartWait::Outgoing(Some(frame)) => buffered_outgoing.push_back(frame), + SseStartWait::Outgoing(Some(frame)) => { + if buffered_outgoing.len() + posts.len() >= self.max_tasks { + return Err(post_capacity_error()); + } + buffered_outgoing.push_back(frame); + } SseStartWait::Outgoing(None) => return Ok(SseStartOutcome::OutgoingClosed), } } @@ -647,16 +770,20 @@ impl HttpTransportLifecycle { fn begin_sse( &mut self, session_id: Option, - event_tx: UnboundedSender, - ) -> futures::channel::oneshot::Receiver<()> { + event_tx: mpsc::Sender, + ) -> Result, AcpError> { + if self.sse_tasks.len() >= self.max_tasks { + return Err(AcpError::internal_error().data("HTTP SSE stream capacity exceeded")); + } let (established_tx, established_rx) = futures::channel::oneshot::channel(); self.sse_tasks.push(run_sse( self.connection.clone(), session_id, event_tx, established_tx, + self.admission.clone(), )); - established_rx + Ok(established_rx) } async fn next_sse_failure(&mut self) -> SseFailure { @@ -674,7 +801,7 @@ enum SseStartWait { Failure(SseFailure), SseEvent(Option), Post(CompletedPost), - Outgoing(Option), + Outgoing(Option), } impl Drop for HttpTransportLifecycle { @@ -687,15 +814,17 @@ impl Drop for HttpTransportLifecycle { fn run_sse( connection: HttpConnection, session_id: Option, - event_tx: UnboundedSender, + event_tx: mpsc::Sender, established_tx: futures::channel::oneshot::Sender<()>, + admission: FrameAdmission, ) -> BoxFuture<'static, SseFailure> { Box::pin(async move { let label = session_id.clone(); - let error = match read_sse(connection, session_id, event_tx, established_tx).await { - Ok(()) => "SSE stream closed".to_string(), - Err(e) => e, - }; + let error = + match read_sse(connection, session_id, event_tx, established_tx, admission).await { + Ok(()) => "SSE stream closed".to_string(), + Err(e) => e, + }; warn!(session_id = ?label, "SSE stream ended: {error}"); SseFailure { session_id: label, @@ -710,6 +839,10 @@ struct SseTasks { } impl SseTasks { + fn len(&self) -> usize { + self.handles.len() + } + fn push(&mut self, task: BoxFuture<'static, SseFailure>) { self.handles.push(task); } @@ -732,7 +865,8 @@ struct ClientState { connection: HttpConnection, open_session_streams: HashSet, pending_requests: HashMap>, - incoming: futures::channel::mpsc::UnboundedSender, + pending_request_leases: HashMap>, + incoming: FrameSender, } struct PendingPost { @@ -741,16 +875,18 @@ struct PendingPost { } impl PendingPost { - fn into_completion(self) -> BoxFuture<'static, CompletedPost> { + fn into_completion(self, permit: Option) -> BoxFuture<'static, CompletedPost> { let Self { pending_requests, response, } = self; async move { - CompletedPost { + let completed = CompletedPost { pending_requests, result: response.await, - } + }; + drop(permit); + completed } .boxed() } @@ -764,7 +900,7 @@ struct CompletedPost { #[derive(Default)] struct PostQueue { - queued: VecDeque, + queued: VecDeque<(PendingPost, Option)>, in_flight: Option>, } @@ -778,11 +914,25 @@ impl PostQueues { fn is_empty(&self) -> bool { self.ordered.is_empty() && self.responses.is_empty() } + + fn len(&self) -> usize { + self.ordered.len() + self.responses.len() + } } impl PostQueue { + fn len(&self) -> usize { + self.queued.len() + usize::from(self.in_flight.is_some()) + } + + #[cfg(test)] fn push(&mut self, post: PendingPost) { - self.queued.push_back(post); + self.queued.push_back((post, None)); + self.start_next(); + } + + fn push_budgeted(&mut self, post: PendingPost, permit: FramePermit) { + self.queued.push_back((post, Some(permit))); self.start_next(); } @@ -800,9 +950,9 @@ impl PostQueue { fn start_next(&mut self) { if self.in_flight.is_none() - && let Some(post) = self.queued.pop_front() + && let Some((post, permit)) = self.queued.pop_front() { - self.in_flight = Some(post.into_completion()); + self.in_flight = Some(post.into_completion(permit)); } } @@ -859,14 +1009,14 @@ impl ClientState { message, RawJsonRpcMessage::Response(RpcResponse::Error { .. }) ) { - self.deliver(message); + self.deliver(message).await.map_err(|e| e.to_string())?; self.connection.close().await; return Ok(InitializeOutcome::Rejected); } connection_id .ok_or_else(|| format!("server did not return {HEADER_CONNECTION_ID} header"))?; - self.deliver(message); + self.deliver(message).await.map_err(|e| e.to_string())?; Ok(InitializeOutcome::Connected) } @@ -889,6 +1039,7 @@ impl ClientState { let pending_requests = pending_request_for_message(&msg) .into_iter() .collect::>(); + self.check_pending_request_capacity(pending_requests.len())?; self.track_pending_requests(&pending_requests); let response = async move { @@ -911,6 +1062,7 @@ impl ClientState { frame: TransportFrame, ) -> Result<(PendingPost, Vec), String> { let bookkeeping = FrameBookkeeping::for_frame(&frame)?; + self.check_pending_request_capacity(bookkeeping.pending_requests.len())?; let connection_id = self .connection .connection_id() @@ -952,16 +1104,43 @@ impl ClientState { } } + fn check_pending_request_capacity(&self, additional: usize) -> Result<(), String> { + let limit = self.incoming.admission().limits().max_queued_frames.max(1); + let existing: usize = self.pending_requests.values().map(VecDeque::len).sum(); + if additional > limit.saturating_sub(existing) { + Err("HTTP pending request capacity exceeded".to_string()) + } else { + Ok(()) + } + } + + fn attach_pending_permits( + &mut self, + pending_requests: &[(RequestId, String)], + permit: &FramePermit, + ) { + for (id, _) in pending_requests { + self.pending_request_leases + .entry(id.clone()) + .or_default() + .push_back(permit.clone()); + } + } + fn remove_pending_requests(&mut self, pending_requests: &[(RequestId, String)]) { for (id, method) in pending_requests.iter().rev() { let remove_entry = self.pending_requests.get_mut(id).is_some_and(|methods| { if let Some(index) = methods.iter().rposition(|candidate| candidate == method) { methods.remove(index); + if let Some(leases) = self.pending_request_leases.get_mut(id) { + leases.remove(index); + } } methods.is_empty() }); if remove_entry { self.pending_requests.remove(id); + self.pending_request_leases.remove(id); } } } @@ -971,8 +1150,12 @@ impl ClientState { let methods = self.pending_requests.get_mut(id)?; (methods.pop_front(), methods.is_empty()) }; + if let Some(leases) = self.pending_request_leases.get_mut(id) { + leases.pop_front(); + } if remove_entry { self.pending_requests.remove(id); + self.pending_request_leases.remove(id); } method } @@ -1031,14 +1214,20 @@ impl ClientState { } } - fn deliver(&self, msg: RawJsonRpcMessage) { - self.deliver_frame(TransportFrame::Single(msg)); + async fn deliver(&self, msg: RawJsonRpcMessage) -> Result<(), AcpError> { + self.deliver_frame(TransportFrame::Single(msg)).await } - fn deliver_frame(&self, frame: TransportFrame) { - if self.incoming.unbounded_send(frame).is_err() { - debug!("upstream channel closed; dropping inbound message"); - } + async fn deliver_frame(&self, frame: TransportFrame) -> Result<(), AcpError> { + self.incoming.send_frame(frame).await + } + + async fn deliver_budgeted(&self, frame: BudgetedFrame) -> Result<(), AcpError> { + self.incoming + .clone() + .send(frame) + .await + .map_err(|error| AcpError::internal_error().data(format!("deliver SSE frame: {error}"))) } } @@ -1096,8 +1285,9 @@ fn is_session_opening_method(method: &str) -> bool { async fn read_sse( connection: HttpConnection, session_id: Option, - event_tx: UnboundedSender, + mut event_tx: mpsc::Sender, established_tx: futures::channel::oneshot::Sender<()>, + admission: FrameAdmission, ) -> Result<(), String> { let connection_id = connection .connection_id() @@ -1117,16 +1307,45 @@ async fn read_sse( trace!(session_id = ?session_id, "SSE stream open"); let _ = established_tx.send(()); - let mut events = eventsource_stream::EventStream::new(response.bytes_stream()); + // Cap each event before EventStream buffers its data fields or JSON parsing + // materializes the payload. A blank line terminates one SSE event. + let max_frame_bytes = admission.limits().max_frame_bytes; + let mut event_bytes = 0usize; + let mut line_has_data = false; + let mut events = + eventsource_stream::EventStream::new(response.bytes_stream().map(move |chunk| { + let chunk = chunk.map_err(std::io::Error::other)?; + for &byte in &chunk { + event_bytes += 1; + if event_bytes > max_frame_bytes { + return Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "SSE event exceeds maximum JSON-RPC frame size", + )); + } + if byte == b'\n' { + if !line_has_data { + event_bytes = 0; + } + line_has_data = false; + } else if byte != b'\r' { + line_has_data = true; + } + } + Ok(chunk) + })); while let Some(event) = events.next().await { let event = event.map_err(|e| e.to_string())?; let payload = event.data; if payload.is_empty() { continue; } - let frame = TransportFrame::parse_json(&payload); + let frame = admission + .admit(TransportFrame::parse_json(&payload)) + .await + .map_err(|error| error.to_string())?; - if event_tx.unbounded_send(SseMessage { frame }).is_err() { + if event_tx.send(SseMessage { frame }).await.is_err() { return Err("upstream channel closed".to_string()); } } @@ -1196,7 +1415,7 @@ where } = channel; let writer = async move { while let Some(frame) = outgoing.next().await { - let text = match frame.to_json() { + let text = match frame.frame().to_json() { Ok(text) => text, Err(error) => { error!("failed to serialize outbound frame: {error}"); @@ -1222,7 +1441,7 @@ where continue; } let frame = TransportFrame::parse_json(text.as_str()); - if incoming.unbounded_send(frame).is_err() { + if incoming.send_frame(frame).await.is_err() { debug!( "upstream channel closed; discarding WS input while draining output" ); @@ -1256,6 +1475,10 @@ where } } +#[cfg(test)] +#[path = "client_admission_tests.rs"] +mod admission_tests; + #[cfg(test)] mod tests { use std::{ @@ -1287,9 +1510,7 @@ mod tests { struct PostsThenExitClient { finish: Arc, finished: Arc, - escaped_tx: futures::channel::oneshot::Sender< - futures::channel::mpsc::UnboundedSender, - >, + escaped_tx: futures::channel::oneshot::Sender, } struct InitializeThenExitClient { @@ -1299,7 +1520,7 @@ mod tests { struct QueueOutgoingThenText { text: Option, - outgoing: Option>, + outgoing: Option, } struct RecordingWsSink(mpsc::UnboundedSender); @@ -1339,6 +1560,12 @@ mod tests { } } + impl TransportFrameTestExt for agent_client_protocol::BudgetedFrame { + fn unwrap(self) -> RawJsonRpcMessage { + into_single_message(self.into_frame()).unwrap() + } + } + #[test] fn malformed_response_shapes_bypass_only_when_the_whole_frame_is_response_only() { let standalone_response = TransportFrame::parse_json( @@ -1374,12 +1601,13 @@ mod tests { reqwest::Client::new(), ); connection.set_connection_id("connection-1".to_string()); - let (incoming, _incoming_rx) = mpsc::unbounded(); + let (incoming, _incoming_rx) = Channel::duplex(); ClientState { connection, open_session_streams: HashSet::new(), pending_requests: HashMap::new(), - incoming, + pending_request_leases: HashMap::new(), + incoming: incoming.tx, } } @@ -1472,6 +1700,57 @@ mod tests { ); } + #[test] + fn cancel_ack_keeps_session_opening_context_until_terminal_response() { + for method in ["session/new", "session/fork"] { + let mut state = initialized_client_state(); + let id = RequestId::Number(7); + state.track_pending_requests(&[(id.clone(), method.into())]); + handle_completed_post( + &mut state, + CompletedPost { + pending_requests: Vec::new(), + result: Ok(()), + }, + ) + .unwrap(); + assert_eq!(state.pending_requests.get(&id).unwrap().len(), 1); + let response = single_frame(RawJsonRpcMessage::response( + id.clone(), + Ok(json!({"sessionId": "new-session"})), + )); + assert_eq!( + state.sessions_to_open_for_responses(&response), + ["new-session"] + ); + assert!(state.pending_requests.is_empty()); + assert!(state.sessions_to_open_for_responses(&response).is_empty()); + } + } + + #[test] + fn only_pure_cancellation_batches_bypass_ordered_posts() { + let cancel = || { + RawJsonRpcMessage::notification( + "_proxy/successor".into(), + json!({"method": "$/cancel_request", "params": {"requestId": 7}}), + ) + .unwrap() + }; + assert!(is_cancellation_frame(&single_frame(cancel()))); + let batch = |other| { + TransportFrame::Batch(TransportBatch::from_messages([cancel(), other]).unwrap()) + }; + assert!(is_cancellation_frame(&batch(cancel()))); + assert!(is_cancellation_frame(&batch(RawJsonRpcMessage::response( + RequestId::Number(7), + Ok(json!({})) + )))); + assert!(!is_cancellation_frame(&batch( + RawJsonRpcMessage::notification("custom/data".into(), json!({})).unwrap() + ))); + } + impl WsSink for RecordingWsSink { fn send( &mut self, @@ -1515,7 +1794,7 @@ mod tests { if let Some(outgoing) = self.outgoing.take() { for method in ["custom/first", "custom/second"] { outgoing - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(), )) .unwrap(); @@ -1559,7 +1838,7 @@ mod tests { })?; channel .tx - .unbounded_send(single_frame( + .send_frame(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -1567,19 +1846,28 @@ mod tests { ) .unwrap(), )) + .await .map_err(|e| { AcpError::internal_error().data(format!("send initialize: {e}")) })?; - into_single_message(channel.rx.next().await.ok_or_else(|| { - AcpError::internal_error().data("initialize response channel closed") - })?)?; + into_single_message( + channel + .rx + .next() + .await + .ok_or_else(|| { + AcpError::internal_error().data("initialize response channel closed") + })? + .into_frame(), + )?; for method in ["custom/first", "custom/second"] { channel .tx - .unbounded_send(single_frame( + .send_frame(single_frame( RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(), )) + .await .map_err(|e| { AcpError::internal_error().data(format!("send {method}: {e}")) })?; @@ -1605,7 +1893,7 @@ mod tests { let client = async move { channel .tx - .unbounded_send(single_frame( + .send_frame(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -1613,12 +1901,20 @@ mod tests { ) .unwrap(), )) + .await .map_err(|error| { AcpError::internal_error().data(format!("send initialize: {error}")) })?; - into_single_message(channel.rx.next().await.ok_or_else(|| { - AcpError::internal_error().data("initialize response channel closed") - })?)?; + into_single_message( + channel + .rx + .next() + .await + .ok_or_else(|| { + AcpError::internal_error().data("initialize response channel closed") + })? + .into_frame(), + )?; sse_started.notified().await; finished.notify_one(); @@ -1714,7 +2010,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -1732,7 +2028,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification( "$/cancel_request".to_string(), json!({ @@ -1765,6 +2061,99 @@ mod tests { server.abort(); } + #[tokio::test] + async fn wrapped_cancellation_post_bypasses_blocked_ordered_post_without_session_header() { + let slow_started = Arc::new(Notify::new()); + let release_slow = Arc::new(Notify::new()); + let (cancel_tx, mut cancel_rx) = tokio::sync::mpsc::unbounded_channel(); + let app = Router::new().route( + "/acp", + post({ + let slow_started = slow_started.clone(); + let release_slow = release_slow.clone(); + move |headers: HeaderMap, body: String| { + let slow_started = slow_started.clone(); + let release_slow = release_slow.clone(); + let cancel_tx = cancel_tx.clone(); + async move { + let value: serde_json::Value = serde_json::from_str(&body).unwrap(); + match value["method"].as_str() { + Some("initialize") => initialize_response().await.into_response(), + Some("session/load") => { + slow_started.notify_one(); + release_slow.notified().await; + StatusCode::ACCEPTED.into_response() + } + Some("_proxy/successor") => { + cancel_tx.send((headers, value)).unwrap(); + StatusCode::ACCEPTED.into_response() + } + other => panic!("unexpected POST: {other:?}"), + } + } + } + }) + .get(pending_sse) + .delete(|| async { StatusCode::ACCEPTED }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let (mut caller, transport) = Channel::duplex(); + let transport = tokio::spawn(run( + HttpClient::new(format!("http://{addr}")).unwrap(), + transport, + )); + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::request("initialize".into(), json!({}), RequestId::Number(1)) + .unwrap(), + )) + .unwrap(); + timeout(Duration::from_secs(1), caller.rx.next()) + .await + .unwrap() + .unwrap(); + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::request( + "session/load".into(), + json!({"sessionId": "source"}), + RequestId::Number(2), + ) + .unwrap(), + )) + .unwrap(); + timeout(Duration::from_secs(1), slow_started.notified()) + .await + .unwrap(); + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::notification( + "_proxy/successor".into(), + json!({ + "method": "$/cancel_request", + "params": {"requestId": 2, "sessionId": "source"} + }), + ) + .unwrap(), + )) + .unwrap(); + let (headers, cancellation) = timeout(Duration::from_secs(1), cancel_rx.recv()) + .await + .expect("cancellation must bypass an in-flight session POST") + .unwrap(); + assert!(headers.get(HEADER_SESSION_ID).is_none()); + assert_eq!(cancellation["params"]["method"], "$/cancel_request"); + release_slow.notify_one(); + transport.abort(); + drop(caller); + server.abort(); + } + #[tokio::test] async fn http_preserves_batch_frames_across_post_and_sse() { let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel(); @@ -1832,7 +2221,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -1862,7 +2251,7 @@ mod tests { ]); caller .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ RawJsonRpcMessage::notification("custom/outbound-one".to_string(), json!({})) .unwrap(), @@ -1884,11 +2273,12 @@ mod tests { .await .unwrap() .unwrap(); - assert!(matches!(&inbound, TransportFrame::Batch(_))); + assert!(matches!(inbound.frame(), TransportFrame::Batch(_))); assert_eq!( - serde_json::from_str::(&inbound.to_json().unwrap()).unwrap(), + serde_json::from_str::(&inbound.frame().to_json().unwrap()).unwrap(), inbound_batch ); + drop(inbound); drop(caller); timeout(Duration::from_secs(1), transport) @@ -1907,7 +2297,6 @@ mod tests { let post_count = Arc::new(AtomicUsize::new(0)); let emit_response = Arc::new(Notify::new()); let connection_stream_established = Arc::new(AtomicBool::new(false)); - let source_stream_established = Arc::new(AtomicBool::new(false)); let response_batch = json!([ { "jsonrpc": "2.0", @@ -1920,20 +2309,16 @@ mod tests { post({ let post_count = post_count.clone(); let connection_stream_established = connection_stream_established.clone(); - let source_stream_established = source_stream_established.clone(); move |body: String| { let post_count = post_count.clone(); let post_tx = post_tx.clone(); let connection_stream_established = connection_stream_established.clone(); - let source_stream_established = source_stream_established.clone(); async move { if post_count.fetch_add(1, Ordering::SeqCst) == 0 { return initialize_response().await.into_response(); } - if !connection_stream_established.load(Ordering::SeqCst) - || !source_stream_established.load(Ordering::SeqCst) - { + if !connection_stream_established.load(Ordering::SeqCst) { return StatusCode::CONFLICT.into_response(); } post_tx @@ -1947,13 +2332,13 @@ mod tests { let emit_response = emit_response.clone(); let response_batch = response_batch.clone(); let connection_stream_established = connection_stream_established.clone(); - let source_stream_established = source_stream_established.clone(); + let post_count = post_count.clone(); move |headers: HeaderMap| { let emit_response = emit_response.clone(); let response_batch = response_batch.clone(); let get_tx = get_tx.clone(); let connection_stream_established = connection_stream_established.clone(); - let source_stream_established = source_stream_established.clone(); + let post_count = post_count.clone(); async move { let session_id = headers .get(HEADER_SESSION_ID) @@ -1961,13 +2346,15 @@ mod tests { .map(String::from); let is_connection_stream = session_id.is_none(); let is_source_stream = session_id.as_deref() == Some("source-session"); + if is_source_stream && post_count.load(Ordering::SeqCst) < 2 { + return StatusCode::CONFLICT.into_response(); + } if is_connection_stream { sleep(Duration::from_millis(50)).await; connection_stream_established.store(true, Ordering::SeqCst); } if is_source_stream { sleep(Duration::from_millis(50)).await; - source_stream_established.store(true, Ordering::SeqCst); } get_tx.send(session_id).unwrap(); @@ -1980,7 +2367,7 @@ mod tests { } futures::future::pending::<()>().await; }; - Sse::new(stream) + Sse::new(stream).into_response() } } }) @@ -1997,7 +2384,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2013,7 +2400,7 @@ mod tests { caller .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([RawJsonRpcMessage::request( "session/fork".to_string(), json!({ "sessionId": "source-session" }), @@ -2040,16 +2427,49 @@ mod tests { .unwrap(); assert!(posted.is_array(), "outgoing batch must remain an array"); + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::notification("$/cancel_request".into(), json!({"requestId": 2})) + .unwrap(), + )) + .unwrap(); + let cancellation = timeout(Duration::from_secs(1), post_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(cancellation["method"], "$/cancel_request"); + assert_eq!(cancellation["params"]["requestId"], 2); + // Control POSTs are serialized with each other. Observing a second + // (unknown-ID, harmless) cancellation proves the client processed the + // first POST's 202 before the original successful response is emitted. + caller + .tx + .try_send(single_frame( + RawJsonRpcMessage::notification( + "$/cancel_request".into(), + json!({"requestId": 999}), + ) + .unwrap(), + )) + .unwrap(); + let barrier = timeout(Duration::from_secs(1), post_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(barrier["params"]["requestId"], 999); emit_response.notify_one(); let response = timeout(Duration::from_secs(1), caller.rx.next()) .await .unwrap() .unwrap(); - assert!(matches!(&response, TransportFrame::Batch(_))); + assert!(matches!(response.frame(), TransportFrame::Batch(_))); assert_eq!( - serde_json::from_str::(&response.to_json().unwrap()).unwrap(), + serde_json::from_str::(&response.frame().to_json().unwrap()) + .unwrap(), response_batch ); + drop(response); let forked_stream = timeout(Duration::from_secs(1), get_rx.recv()) .await .unwrap() @@ -2129,7 +2549,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2153,7 +2573,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "custom/sessionish".to_string(), json!({}), @@ -2258,7 +2678,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2282,7 +2702,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "session/fork".to_string(), json!({ "sessionId": "source-session" }), @@ -2381,7 +2801,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2397,7 +2817,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification("custom/slow".to_string(), json!({})).unwrap(), )) .unwrap(); @@ -2407,7 +2827,7 @@ mod tests { caller .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ RawJsonRpcMessage::notification("custom/one".to_string(), json!({})).unwrap(), RawJsonRpcMessage::notification("custom/two".to_string(), json!({})).unwrap(), @@ -2424,7 +2844,7 @@ mod tests { caller .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ RawJsonRpcMessage::response(RequestId::Number(10), Ok(json!({}))), RawJsonRpcMessage::response(RequestId::Number(11), Ok(json!({}))), @@ -2529,7 +2949,7 @@ mod tests { ); assert!( escaped - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification("custom/too-late".to_string(), json!({}),) .unwrap() )) @@ -2622,12 +3042,13 @@ mod tests { reqwest::Client::new(), ); connection.set_connection_id("connection-1".to_string()); - let (incoming, _incoming_rx) = mpsc::unbounded(); + let (incoming, _incoming_rx) = Channel::duplex(); let mut state = ClientState { connection: connection.clone(), open_session_streams: HashSet::new(), pending_requests: HashMap::new(), - incoming, + pending_request_leases: HashMap::new(), + incoming: incoming.tx, }; let pending_request = (RequestId::Number(7), "custom/earlier".to_string()); state.track_pending_requests(std::slice::from_ref(&pending_request)); @@ -2637,15 +3058,19 @@ mod tests { response: async { Err("earlier post failed".to_string()) }.boxed(), }); - let (_outgoing_tx, mut outgoing) = mpsc::unbounded(); + let (outgoing_channel, _outgoing_peer) = Channel::duplex(); + let mut outgoing = outgoing_channel.rx; let mut buffered_outgoing = VecDeque::new(); - let (event_tx, mut event_rx) = mpsc::unbounded(); - let mut lifecycle = HttpTransportLifecycle::new(connection); + let (event_tx, mut event_rx) = mpsc::channel(16); + let admission = state.incoming.admission(); + let max_tasks = admission.limits().max_queued_frames; + let mut lifecycle = HttpTransportLifecycle::new(connection, admission, max_tasks); let error = timeout( Duration::from_secs(1), lifecycle.start_sse( Some("later-session".to_string()), event_tx, + true, SseStartContext { events: &mut event_rx, outgoing: &mut outgoing, @@ -2668,6 +3093,15 @@ mod tests { #[tokio::test] async fn stalled_sse_establishment_keeps_callback_responses_moving() { + check_callback_response_progress(false).await; + } + + #[tokio::test] + async fn cold_session_post_admission_keeps_callback_responses_moving() { + check_callback_response_progress(true).await; + } + + async fn check_callback_response_progress(wait_for_post: bool) { let release_get = Arc::new(Notify::new()); let complete_earlier_post = Arc::new(Notify::new()); let app = Router::new().route( @@ -2703,12 +3137,13 @@ mod tests { reqwest::Client::new(), ); connection.set_connection_id("connection-1".to_string()); - let (incoming, mut incoming_rx) = mpsc::unbounded(); + let (incoming, mut incoming_peer) = Channel::duplex(); let mut state = ClientState { connection: connection.clone(), open_session_streams: HashSet::new(), pending_requests: HashMap::new(), - incoming, + pending_request_leases: HashMap::new(), + incoming: incoming.tx, }; let mut posts = PostQueues::default(); posts.ordered.push(PendingPost { @@ -2720,46 +3155,56 @@ mod tests { .boxed(), }); - let (outgoing_tx, mut outgoing) = mpsc::unbounded(); + let (outgoing_channel, outgoing_peer) = Channel::duplex(); + let outgoing_tx = outgoing_peer.tx; + let mut outgoing = outgoing_channel.rx; let outgoing_guard = outgoing_tx.clone(); let mut buffered_outgoing = VecDeque::new(); - let (event_tx, mut event_rx) = mpsc::unbounded(); + let (mut event_tx, mut event_rx) = mpsc::channel(16); event_tx - .unbounded_send(SseMessage { - frame: single_frame( - RawJsonRpcMessage::request( - "test/callback".to_string(), - json!({}), - RequestId::Number(99), - ) + .try_send(SseMessage { + frame: state + .incoming + .admission() + .try_admit(single_frame( + RawJsonRpcMessage::request( + "test/callback".to_string(), + json!({}), + RequestId::Number(99), + ) + .unwrap(), + )) .unwrap(), - ), }) .unwrap(); let responder = async move { - let callback = incoming_rx + let callback = incoming_peer + .rx .next() .await .expect("callback was not delivered"); assert!(matches!( - into_single_message(callback).unwrap(), + into_single_message(callback.into_frame()).unwrap(), RawJsonRpcMessage::Request(request) if request.method.as_ref() == "test/callback" )); outgoing_tx - .unbounded_send(single_frame(RawJsonRpcMessage::response( + .try_send(single_frame(RawJsonRpcMessage::response( RequestId::Number(99), Ok(json!({})), ))) .unwrap(); }; - let mut lifecycle = HttpTransportLifecycle::new(connection); + let admission = state.incoming.admission(); + let max_tasks = admission.limits().max_queued_frames; + let mut lifecycle = HttpTransportLifecycle::new(connection, admission, max_tasks); let (outcome, ()) = timeout(Duration::from_secs(1), async { futures::join!( lifecycle.start_sse( Some("later-session".to_string()), event_tx, + wait_for_post, SseStartContext { events: &mut event_rx, outgoing: &mut outgoing, @@ -2821,7 +3266,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2840,7 +3285,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(), )) .unwrap(); @@ -2942,7 +3387,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -2963,7 +3408,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "custom/slow".to_string(), json!({}), @@ -2987,7 +3432,7 @@ mod tests { caller .tx - .unbounded_send(single_frame(RawJsonRpcMessage::response( + .try_send(single_frame(RawJsonRpcMessage::response( RequestId::Number(99), Ok(json!({})), ))) @@ -3039,7 +3484,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3057,7 +3502,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "session/prompt".to_string(), json!({}), @@ -3103,7 +3548,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3158,7 +3603,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3179,7 +3624,7 @@ mod tests { .unwrap() .unwrap(); - let TransportFrame::Malformed { raw, error } = frame else { + let TransportFrame::Malformed { raw, error } = frame.frame() else { panic!("expected malformed frame, got {frame:?}"); }; assert_eq!(raw, "{not json"); @@ -3211,7 +3656,7 @@ mod tests { .await .unwrap() .unwrap(); - let TransportFrame::Malformed { raw, error } = frame else { + let TransportFrame::Malformed { raw, error } = frame.frame() else { panic!("expected malformed frame, got {frame:?}"); }; assert_eq!(raw, "{not json"); @@ -3243,7 +3688,7 @@ mod tests { } = caller; drop(incoming); outgoing - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ RawJsonRpcMessage::notification("custom/first".to_string(), json!({})).unwrap(), RawJsonRpcMessage::notification("custom/second".to_string(), json!({})) @@ -3335,7 +3780,7 @@ mod tests { } = caller; drop(incoming); outgoing - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(), )) .unwrap(); @@ -3416,7 +3861,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3474,7 +3919,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3525,7 +3970,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), @@ -3584,7 +4029,7 @@ mod tests { caller .tx - .unbounded_send(single_frame( + .try_send(single_frame( RawJsonRpcMessage::request( "initialize".to_string(), json!({}), diff --git a/src/agent-client-protocol-http/src/client_admission_tests.rs b/src/agent-client-protocol-http/src/client_admission_tests.rs new file mode 100644 index 00000000..58dfa374 --- /dev/null +++ b/src/agent-client-protocol-http/src/client_admission_tests.rs @@ -0,0 +1,301 @@ +use std::{convert::Infallible, time::Duration}; + +use agent_client_protocol::ConnectionLimits; +use axum::{ + Router, + response::{Sse, sse::Event}, + routing::{delete, get}, +}; +use futures::{StreamExt, channel::mpsc}; +use tokio::{net::TcpListener, time::timeout}; + +use super::*; + +#[tokio::test] +async fn sse_staging_holds_shared_budget_until_frame_is_released() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".to_string(), serde_json::json!({})).unwrap(), + ); + let json = frame.to_json().unwrap(); + let frame_bytes = json.len(); + let limits = ConnectionLimits { + max_frame_bytes: frame_bytes + 128, + // Leave room for one data event and reserve a whole frame for control. + max_queued_bytes: frame_bytes + frame_bytes + 128, + max_queued_frames: 4, + }; + let (_caller, transport) = Channel::duplex_with_limits(limits); + let admission = transport.tx.admission(); + let app = Router::new().route( + "/acp", + get({ + let json = json.clone(); + move || { + let json = json.clone(); + async move { + Sse::new(futures::stream::iter((0..3).map(move |_| { + Ok::<_, Infallible>(Event::default().data(json.clone())) + }))) + } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let connection = HttpConnection::new( + url::Url::parse(&format!("http://{address}/acp")).unwrap(), + reqwest::Client::new(), + ); + connection.set_connection_id("test-connection".into()); + let (event_tx, mut event_rx) = mpsc::channel(4); + let (established_tx, established_rx) = futures::channel::oneshot::channel(); + let reader = tokio::spawn(read_sse( + connection, + None, + event_tx, + established_tx, + admission.clone(), + )); + timeout(Duration::from_secs(2), established_rx) + .await + .unwrap() + .unwrap(); + let first = timeout(Duration::from_secs(2), event_rx.next()) + .await + .unwrap() + .unwrap(); + assert_eq!(first.frame.frame().to_json().unwrap(), json); + assert!(admission.try_admit(frame.clone()).is_err()); + // A second event may be parsed, but cannot enter the staging queue until + // the first event's shared charge is released. + assert!( + timeout(Duration::from_millis(40), event_rx.next()) + .await + .is_err() + ); + drop(first); + let second = timeout(Duration::from_secs(2), event_rx.next()) + .await + .unwrap() + .unwrap(); + assert!(admission.try_admit(frame.clone()).is_err()); + reader.abort(); + // Dropping the SSE reader releases even an event admitted but still + // waiting to send; dropping the receiver releases queued events too. + drop(second); + drop(event_rx); + reader.await.unwrap_err(); + let recovered = admission + .try_admit(frame) + .expect("cancelled SSE released permits"); + drop(recovered); + server.abort(); +} + +#[tokio::test] +async fn post_and_stream_counts_are_bounded_independently_of_frame_bytes() { + let cancellation = TransportFrame::Single( + RawJsonRpcMessage::notification( + "$/cancel_request".to_string(), + serde_json::json!({"requestId": 1}), + ) + .unwrap(), + ); + assert!(is_cancellation_frame(&cancellation)); + assert!(!is_response_only_frame(&cancellation)); + let (_caller, transport) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 4096, + max_queued_frames: 3, + }); + let connection = HttpConnection::new( + url::Url::parse("http://127.0.0.1:1/acp").unwrap(), + reqwest::Client::new(), + ); + let mut lifecycle = HttpTransportLifecycle::new(connection, transport.tx.admission(), 3); + for index in 0..3 { + drop( + lifecycle + .begin_sse(Some(index.to_string()), mpsc::channel(1).0) + .unwrap(), + ); + } + assert!( + lifecycle + .begin_sse(Some("excess".into()), mpsc::channel(1).0) + .is_err() + ); + lifecycle.sse_tasks.abort_all(); + assert_eq!(lifecycle.sse_tasks.len(), 0); + + let mut posts = PostQueues::default(); + for _ in 0..2 { + check_post_capacity(&posts, 3, false).unwrap(); + posts.ordered.push(PendingPost { + pending_requests: Vec::new(), + response: Box::pin(futures::future::pending()), + }); + } + assert!(check_post_capacity(&posts, 3, false).is_err()); + check_post_capacity(&posts, 3, true).unwrap(); + posts.responses.push(PendingPost { + pending_requests: Vec::new(), + response: Box::pin(futures::future::pending()), + }); + assert_eq!(posts.len(), 3); + assert!(check_post_capacity(&posts, 3, true).is_err()); + drop(posts); +} + +#[tokio::test] +async fn cancelled_post_releases_its_body_budget() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".to_string(), serde_json::json!({})).unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let (_caller, transport) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes + 128, + max_queued_bytes: bytes * 2 + 128, + max_queued_frames: 4, + }); + let admission = transport.tx.admission(); + let budgeted = admission.try_admit(frame.clone()).unwrap(); + let (_, permit) = budgeted.into_parts(); + let mut posts = PostQueues::default(); + posts.ordered.push_budgeted( + PendingPost { + pending_requests: Vec::new(), + response: Box::pin(futures::future::pending()), + }, + permit, + ); + assert!(admission.try_admit(frame.clone()).is_err()); + drop(posts); + assert!(admission.try_admit(frame).is_ok()); +} + +#[tokio::test] +async fn delivering_sse_preserves_admission_through_output_channel() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".to_string(), serde_json::json!({})).unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let (mut caller, transport) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes + 128, + max_queued_bytes: bytes * 2 + 128, + max_queued_frames: 4, + }); + let admission = transport.tx.admission(); + let state = ClientState { + connection: HttpConnection::new( + url::Url::parse("http://127.0.0.1:1/acp").unwrap(), + reqwest::Client::new(), + ), + open_session_streams: HashSet::new(), + pending_requests: HashMap::new(), + pending_request_leases: HashMap::new(), + incoming: transport.tx, + }; + state + .deliver_budgeted(admission.try_admit(frame.clone()).unwrap()) + .await + .unwrap(); + assert!(admission.try_admit(frame.clone()).is_err()); + let delivered = caller.rx.next().await.unwrap(); + assert!(admission.try_admit(frame.clone()).is_err()); + drop(delivered); + assert!(admission.try_admit(frame).is_ok()); +} + +#[tokio::test] +async fn pending_requests_hold_their_source_charge_until_terminal_response() { + let request = RawJsonRpcMessage::request( + "test/request".to_string(), + serde_json::json!({}), + RequestId::Number(1), + ) + .unwrap(); + let frame = TransportFrame::Single(request); + let bytes = frame.to_json().unwrap().len(); + let (_caller, transport) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes + 128, + max_queued_bytes: bytes * 2 + 128, + max_queued_frames: 1, + }); + let admission = transport.tx.admission(); + let mut state = ClientState { + connection: HttpConnection::new( + url::Url::parse("http://127.0.0.1:1/acp").unwrap(), + reqwest::Client::new(), + ), + open_session_streams: HashSet::new(), + pending_requests: HashMap::new(), + pending_request_leases: HashMap::new(), + incoming: transport.tx, + }; + state.connection.set_connection_id("connection-1".into()); + let (_, permit) = admission.try_admit(frame.clone()).unwrap().into_parts(); + let post = state.prepare_frame_post(frame.clone()).unwrap().0; + state.attach_pending_permits(&post.pending_requests, &permit); + assert!(state.check_pending_request_capacity(1).is_err()); + drop(post); + drop(permit); + assert!(admission.try_admit(frame.clone()).is_err()); + assert_eq!( + state + .take_pending_request_method(&RequestId::Number(1)) + .as_deref(), + Some("test/request") + ); + assert!(state.check_pending_request_capacity(1).is_ok()); + let (_, permit) = admission.try_admit(frame.clone()).unwrap().into_parts(); + let post = state.prepare_frame_post(frame.clone()).unwrap().0; + state.attach_pending_permits(&post.pending_requests, &permit); + drop(post); + drop(permit); + let cancel = RawJsonRpcMessage::notification( + "$/cancel_request".into(), + serde_json::json!({"requestId": 1}), + ) + .unwrap(); + let post = state.prepare_post(cancel).unwrap(); + handle_completed_post( + &mut state, + CompletedPost { + pending_requests: post.pending_requests, + result: Ok(()), + }, + ) + .unwrap(); + assert!(state.check_pending_request_capacity(1).is_err()); + assert!(admission.try_admit(frame.clone()).is_err()); + assert_eq!( + state + .take_pending_request_method(&RequestId::Number(1)) + .as_deref(), + Some("test/request") + ); + assert!(state.check_pending_request_capacity(1).is_ok()); + assert!(admission.try_admit(frame).is_ok()); +} + +#[tokio::test] +async fn unresponsive_close_does_not_stall_transport_shutdown() { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let app = Router::new().route( + "/acp", + delete(|| async { futures::future::pending::().await }), + ); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let connection = HttpConnection::new( + url::Url::parse(&format!("http://{address}/acp")).unwrap(), + reqwest::Client::new(), + ); + connection.set_connection_id("test-connection".into()); + timeout(Duration::from_secs(4), connection.close()) + .await + .expect("DELETE must not indefinitely block transport shutdown"); + server.abort(); +} diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index 3fdbe0f5..cc1c7d68 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -4,7 +4,8 @@ use std::{ }; use agent_client_protocol::{ - Channel, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame, + BudgetedFrame, Channel, ConnectionLimits, FrameAdmission, FramePermit, RawJsonRpcMessage, + RawJsonRpcResponse as RpcResponse, TransportBatch, TransportBatchEntry, TransportFrame, schema::v1::RequestId, }; use futures::{SinkExt, StreamExt}; @@ -20,14 +21,15 @@ pub(crate) enum ResponseRoute { } enum OutboundTransport { - Http(HttpOutbound), + Http(Box), WebSocket(WebSocketOutbound), } struct HttpOutbound { connection_stream: OutboundMailbox, - session_streams: RwLock>>, - pending_routes: Mutex>>, + session_streams: RwLock, Option)>>, + pending_routes: Mutex)>>>, + limits: ConnectionLimits, } struct WebSocketOutbound { @@ -35,28 +37,43 @@ struct WebSocketOutbound { } struct OutboundMailbox { - sender: mpsc::UnboundedSender, - receiver_slot: Arc>>>, + sender: mpsc::Sender, + receiver_slot: Arc>>>, +} + +struct OutboundValue { + text: String, + permit: Option, } pub(crate) struct OutboundLease { - receiver: Option>, - receiver_slot: Arc>>>, + receiver: Option>, + receiver_slot: Arc>>>, + current: Option, } impl OutboundMailbox { fn new() -> Self { - let (sender, receiver) = mpsc::unbounded_channel(); + let (sender, receiver) = mpsc::channel(32); Self { sender, receiver_slot: Arc::new(StdMutex::new(Some(receiver))), } } + #[cfg(test)] fn push(&self, msg: String) -> Result<(), &'static str> { + self.push_with_permit(msg, None) + } + + fn push_with_permit( + &self, + text: String, + permit: Option, + ) -> Result<(), &'static str> { self.sender - .send(msg) - .map_err(|_| "outbound mailbox receiver closed") + .try_send(OutboundValue { text, permit }) + .map_err(|_| "outbound mailbox full or receiver closed") } fn try_acquire(&self) -> Option { @@ -68,24 +85,35 @@ impl OutboundMailbox { Some(OutboundLease { receiver: Some(receiver), receiver_slot: self.receiver_slot.clone(), + current: None, }) } } impl OutboundLease { pub(crate) async fn recv(&mut self) -> Option { - self.receiver + // The previously returned text has been handed to the transport. Do + // not retain its byte charge while waiting for the next frame. + self.current.take(); + let value = self + .receiver .as_mut() .expect("outbound lease receiver missing") .recv() - .await + .await?; + self.current = value.permit; + Some(value.text) } pub(crate) fn try_recv(&mut self) -> Result { - self.receiver + self.current.take(); + let value = self + .receiver .as_mut() .expect("outbound lease receiver missing") - .try_recv() + .try_recv()?; + self.current = value.permit; + Ok(value.text) } } @@ -104,8 +132,9 @@ impl Drop for OutboundLease { } pub(crate) struct Connection { - inbound_tx: mpsc::UnboundedSender, - outbound_rx: Mutex>>, + inbound_tx: mpsc::Sender, + inbound_admission: FrameAdmission, + outbound_rx: Mutex>>, agent_handle: Mutex>>, router_handle: Mutex>>, closed_tx: watch::Sender, @@ -114,17 +143,53 @@ pub(crate) struct Connection { impl Connection { pub(crate) fn send_frame_to_agent(&self, frame: TransportFrame) -> Result<(), &'static str> { + let frame = self.admit_frame_to_agent(frame)?; + self.send_budgeted_frame_to_agent(frame) + } + + pub(crate) fn agent_channel_closed(&self) -> bool { + self.inbound_tx.is_closed() + } + + pub(crate) fn admit_frame_to_agent( + &self, + frame: TransportFrame, + ) -> Result { + self.inbound_admission + .try_admit(frame) + .map_err(|_| "agent frame byte capacity exceeded") + } + + pub(crate) fn send_budgeted_frame_to_agent( + &self, + frame: BudgetedFrame, + ) -> Result<(), &'static str> { self.inbound_tx - .send(frame) - .map_err(|_| "agent channel closed") + .try_send(frame) + .map_err(|_| "agent channel full or closed") } - pub(crate) async fn record_pending_route(&self, id: RequestId, route: ResponseRoute) { - self.outbound_transport - .record_pending_route(id, route) - .await; + pub(crate) fn reserve_inbound(&self) -> Result, &'static str> { + self.inbound_tx + .clone() + .try_reserve_owned() + .map_err(|_| "agent channel full or closed") } + pub(crate) async fn register_post_routes( + &self, + sessions: &[String], + routes: &[(RequestId, ResponseRoute)], + permit: &FramePermit, + ) -> Result<(), &'static str> { + if let OutboundTransport::Http(http) = &self.outbound_transport { + http.register_post_routes(sessions, routes, permit).await + } else { + Ok(()) + } + } + + #[cfg(test)] pub(crate) async fn ensure_session(&self, session_id: &str) { self.outbound_transport.ensure_session(session_id).await; } @@ -177,11 +242,14 @@ impl Connection { })); } - pub(crate) async fn route_outbound(&self, frame: TransportFrame) -> Result<(), &'static str> { - self.outbound_transport.route_outbound(frame).await + pub(crate) async fn route_outbound(&self, frame: BudgetedFrame) -> Result<(), &'static str> { + let (frame, permit) = frame.into_parts(); + self.outbound_transport + .route_outbound(frame, Some(permit)) + .await } - pub(crate) async fn recv_initial(&self) -> Option { + pub(crate) async fn recv_initial(&self) -> Option { let mut guard = self.outbound_rx.lock().await; let rx = guard.as_mut()?; rx.recv().await @@ -191,6 +259,10 @@ impl Connection { // Explicit peer teardown is abortive. Natural agent completion instead // awaits the router in `close_connection_task` before closing streams. self.close_streams(); + if let OutboundTransport::Http(http) = &self.outbound_transport { + http.session_streams.write().await.clear(); + http.pending_routes.lock().await.clear(); + } if let Some(h) = self.agent_handle.lock().await.take() { h.abort(); } @@ -206,21 +278,14 @@ impl Connection { impl OutboundTransport { fn http() -> Self { - Self::Http(HttpOutbound::new()) + Self::Http(Box::new(HttpOutbound::new())) } fn websocket() -> Self { Self::WebSocket(WebSocketOutbound::new()) } - async fn record_pending_route(&self, id: RequestId, route: ResponseRoute) { - let Self::Http(http) = self else { - return; - }; - - http.record_pending_route(id, route).await; - } - + #[cfg(test)] async fn ensure_session(&self, session_id: &str) { let Self::Http(http) = self else { return; @@ -238,7 +303,12 @@ impl OutboundTransport { async fn subscribe_session_stream(&self, session_id: &str) -> Option { match self { - Self::Http(http) => http.session_stream(session_id).await.try_acquire(), + Self::Http(http) => http + .session_streams + .read() + .await + .get(session_id) + .and_then(|(stream, _)| stream.try_acquire()), Self::WebSocket(_) => None, } } @@ -259,7 +329,11 @@ impl OutboundTransport { http.connection_stream.push(msg) } - async fn route_outbound(&self, frame: TransportFrame) -> Result<(), &'static str> { + async fn route_outbound( + &self, + frame: TransportFrame, + permit: Option, + ) -> Result<(), &'static str> { match frame { TransportFrame::Single(message) => { let serialized = match serde_json::to_string(&message) { @@ -270,13 +344,18 @@ impl OutboundTransport { } }; match self { - Self::Http(http) => http.route_outbound(&message, serialized).await, - Self::WebSocket(websocket) => websocket.all_outbound.push(serialized), + Self::Http(http) => { + http.route_outbound_with_permit(&message, serialized, permit) + .await + } + Self::WebSocket(websocket) => { + websocket.all_outbound.push_with_permit(serialized, permit) + } } } TransportFrame::Malformed { raw, .. } => match self { - Self::Http(http) => http.connection_stream.push(raw), - Self::WebSocket(websocket) => websocket.all_outbound.push(raw), + Self::Http(http) => http.connection_stream.push_with_permit(raw, permit), + Self::WebSocket(websocket) => websocket.all_outbound.push_with_permit(raw, permit), }, TransportFrame::Batch(batch) => { let serialized = match serde_json::to_string(&batch) { @@ -287,8 +366,10 @@ impl OutboundTransport { } }; match self { - Self::Http(http) => http.route_outbound_batch(&batch, serialized).await, - Self::WebSocket(websocket) => websocket.all_outbound.push(serialized), + Self::Http(http) => http.route_outbound_batch(&batch, serialized, permit).await, + Self::WebSocket(websocket) => { + websocket.all_outbound.push_with_permit(serialized, permit) + } } } } @@ -301,9 +382,58 @@ impl HttpOutbound { connection_stream: OutboundMailbox::new(), session_streams: RwLock::new(HashMap::new()), pending_routes: Mutex::new(HashMap::new()), + limits: ConnectionLimits::default(), } } + async fn register_post_routes( + &self, + sessions: &[String], + routes: &[(RequestId, ResponseRoute)], + permit: &FramePermit, + ) -> Result<(), &'static str> { + // Lock both metadata tables in one order and check the whole batch + // before inserting anything: rejection must never leave half a batch. + let mut streams = self.session_streams.write().await; + let mut pending = self.pending_routes.lock().await; + let mut new_sessions = Vec::new(); + for id in sessions { + if !streams.contains_key(id) && !new_sessions.contains(id) { + new_sessions.push(id.clone()); + } + } + let pending_count: usize = pending.values().map(VecDeque::len).sum(); + let limit = self.limits.max_queued_frames.max(1); + let available = limit.saturating_sub(streams.len().saturating_add(pending_count)); + if new_sessions.len().saturating_add(routes.len()) > available { + return Err("HTTP pending route or session capacity exceeded"); + } + // Reserve every new ID before publishing any metadata. Session + // mailboxes retain only their key, not the unrelated POST payload; + // pending response routes still retain their originating frame. + let session_permits = new_sessions + .iter() + .map(|id| { + permit + .try_reserve_metadata(session_metadata_bytes(id)) + .map_err(|_| "HTTP session metadata capacity exceeded") + }) + .collect::, _>>()?; + for (id, session_permit) in new_sessions.into_iter().zip(session_permits) { + streams.insert(id, (Arc::new(OutboundMailbox::new()), Some(session_permit))); + } + for (id, route) in routes { + if let Some(id) = pending_route_key(id) { + pending + .entry(id) + .or_default() + .push_back((route.clone(), Some(permit.clone()))); + } + } + Ok(()) + } + + #[cfg(test)] async fn record_pending_route(&self, id: RequestId, route: ResponseRoute) { if let Some(key) = pending_route_key(&id) { self.pending_routes @@ -311,31 +441,77 @@ impl HttpOutbound { .await .entry(key) .or_default() - .push_back(route); + .push_back((route, None)); } } + #[cfg(test)] async fn ensure_session(&self, session_id: &str) { self.session_stream(session_id).await; } + async fn session_stream_with_permit( + &self, + session_id: &str, + permit: Option, + ) -> Result, &'static str> { + let mut streams = self.session_streams.write().await; + if let Some((stream, _)) = streams.get(session_id) { + return Ok(stream.clone()); + } + let Some(permit) = permit else { + return Err("session stream has no admitted source frame"); + }; + let pending_count: usize = self + .pending_routes + .lock() + .await + .values() + .map(VecDeque::len) + .sum(); + if streams.len().saturating_add(pending_count) >= self.limits.max_queued_frames.max(1) { + return Err("HTTP session stream capacity exceeded"); + } + let session_permit = permit + .try_reserve_metadata(session_metadata_bytes(session_id)) + .map_err(|_| "HTTP session metadata capacity exceeded")?; + let stream = Arc::new(OutboundMailbox::new()); + streams.insert( + session_id.to_string(), + (stream.clone(), Some(session_permit)), + ); + Ok(stream) + } + + #[cfg(test)] async fn session_stream(&self, session_id: &str) -> Arc { if let Some(stream) = self.session_streams.read().await.get(session_id) { - return stream.clone(); + return stream.0.clone(); } self.session_streams .write() .await .entry(session_id.to_string()) - .or_insert_with(|| Arc::new(OutboundMailbox::new())) + .or_insert_with(|| (Arc::new(OutboundMailbox::new()), None)) + .0 .clone() } + #[cfg(test)] async fn route_outbound( &self, msg: &RawJsonRpcMessage, serialized: String, + ) -> Result<(), &'static str> { + self.route_outbound_with_permit(msg, serialized, None).await + } + + async fn route_outbound_with_permit( + &self, + msg: &RawJsonRpcMessage, + serialized: String, + permit: Option, ) -> Result<(), &'static str> { let route = match msg { RawJsonRpcMessage::Request(_) | RawJsonRpcMessage::Notification(_) => { @@ -353,15 +529,23 @@ impl HttpOutbound { route.unwrap_or(ResponseRoute::Connection) } }; + // A successful session/new (or fork) response can be followed + // immediately by a session SSE GET, before any session-scoped POST. + if let Some(session_id) = response_session_id(msg) { + self.session_stream_with_permit(session_id, permit.clone()) + .await?; + } match route { ResponseRoute::Connection => { trace!(target = "connection", "→ connection-scoped stream"); - self.connection_stream.push(serialized) + self.connection_stream.push_with_permit(serialized, permit) } ResponseRoute::Session(sid) => { trace!(target = %sid, "→ session-scoped stream"); - self.session_stream(&sid).await.push(serialized) + self.session_stream_with_permit(&sid, permit.clone()) + .await? + .push_with_permit(serialized, permit) } } } @@ -370,6 +554,7 @@ impl HttpOutbound { &self, batch: &TransportBatch, serialized: String, + permit: Option, ) -> Result<(), &'static str> { let mut pending_routes = self.pending_routes.lock().await; let mut common_route = None; @@ -390,6 +575,14 @@ impl HttpOutbound { } } drop(pending_routes); + for entry in batch.entries() { + if let TransportBatchEntry::Message(message) = entry + && let Some(session_id) = response_session_id(message) + { + self.session_stream_with_permit(session_id, permit.clone()) + .await?; + } + } let route = if routes_disagree { ResponseRoute::Connection @@ -399,11 +592,13 @@ impl HttpOutbound { match route { ResponseRoute::Connection => { trace!(target = "connection", "→ connection-scoped batch stream"); - self.connection_stream.push(serialized) + self.connection_stream.push_with_permit(serialized, permit) } ResponseRoute::Session(session_id) => { trace!(target = %session_id, "→ session-scoped batch stream"); - self.session_stream(&session_id).await.push(serialized) + self.session_stream_with_permit(&session_id, permit.clone()) + .await? + .push_with_permit(serialized, permit) } } } @@ -442,7 +637,8 @@ where Channel, futures::future::BoxFuture<'static, agent_client_protocol::Result<()>>, ) { - self().into_channel_and_future() + let (channel, driver) = self().into_channel_and_future(); + (channel, Box::pin(driver)) } } @@ -483,14 +679,19 @@ impl ConnectionRegistry { outbound_transport: OutboundTransport, ) -> Arc { let (channel, agent_future) = self.factory.spawn_agent(); - let (inbound_tx, mut inbound_rx) = mpsc::unbounded_channel::(); - let (outbound_tx, outbound_rx) = mpsc::unbounded_channel::(); + let mut outbound_transport = outbound_transport; + if let OutboundTransport::Http(http) = &mut outbound_transport { + http.limits = channel.tx.admission().limits(); + } + let (inbound_tx, mut inbound_rx) = mpsc::channel::(32); + let (outbound_tx, outbound_rx) = mpsc::channel::(32); let (closed_tx, _) = watch::channel(false); let Channel { rx: mut agent_rx, tx: mut agent_tx, } = channel; + let inbound_admission = agent_tx.admission(); let inbound = async move { while let Some(msg) = inbound_rx.recv().await { if agent_tx.send(msg).await.is_err() { @@ -502,11 +703,36 @@ impl ConnectionRegistry { let (inbound_abort, inbound_abort_registration) = futures::future::AbortHandle::new_pair(); let inbound = futures::future::Abortable::new(inbound, inbound_abort_registration); let inbound_abort_for_outbound = inbound_abort.clone(); + let mut router_closed = closed_tx.subscribe(); let outbound = async move { - while let Some(msg) = agent_rx.next().await { - if outbound_tx.send(msg).is_err() { - inbound_abort_for_outbound.abort(); - break; + loop { + tokio::select! { + // A fatal router failure must tear down even an idle agent: + // waiting for another frame here can otherwise retain the + // agent and its registry entry forever. + changed = router_closed.changed() => { + if changed.is_err() || *router_closed.borrow() { + inbound_abort_for_outbound.abort(); + break; + } + } + msg = agent_rx.next() => { + let Some(msg) = msg else { break }; + let sent = tokio::select! { + sent = outbound_tx.send(msg) => sent.is_ok(), + changed = router_closed.changed() => { + if changed.is_err() || *router_closed.borrow() { + false + } else { + continue; + } + } + }; + if !sent { + inbound_abort_for_outbound.abort(); + break; + } + } } } }; @@ -516,6 +742,7 @@ impl ConnectionRegistry { let connection = Arc::new(Connection { inbound_tx, + inbound_admission, outbound_rx: Mutex::new(Some(outbound_rx)), agent_handle: Mutex::new(None), router_handle: Mutex::new(None), @@ -582,6 +809,10 @@ async fn close_connection_task(connection: Weak) { error!("outbound router task failed while draining: {error}"); } connection.close_streams(); + if let OutboundTransport::Http(http) = &connection.outbound_transport { + http.session_streams.write().await.clear(); + http.pending_routes.lock().await.clear(); + } } fn pending_route_key(id: &RequestId) -> Option { @@ -591,8 +822,23 @@ fn pending_route_key(id: &RequestId) -> Option { } } +fn session_metadata_bytes(session_id: &str) -> usize { + // The retained key is one JSON string; account for escaping and quotes, + // not for the unrelated source request/response payload. + serde_json::to_string(session_id) + .expect("string serialization cannot fail") + .len() +} + +fn response_session_id(msg: &RawJsonRpcMessage) -> Option<&str> { + let RawJsonRpcMessage::Response(RpcResponse::Result { result, .. }) = msg else { + return None; + }; + result.get("sessionId")?.as_str() +} + fn take_pending_route( - pending_routes: &mut HashMap>, + pending_routes: &mut HashMap)>>, key: &RequestId, ) -> Option { let routes = pending_routes.get_mut(key)?; @@ -601,9 +847,13 @@ fn take_pending_route( if remove_entry { pending_routes.remove(key); } - route + route.map(|(route, _permit)| route) } +#[cfg(test)] +#[path = "connection_admission_tests.rs"] +mod admission_tests; + #[cfg(test)] mod tests { use std::sync::Arc; @@ -617,18 +867,18 @@ mod tests { use super::*; - const ISSUE_288_BURST: usize = 1_025; - #[tokio::test] - async fn outbound_mailbox_buffers_bursts_before_subscription() { + async fn outbound_mailbox_bounds_bursts_before_subscription() { let mailbox = OutboundMailbox::new(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { mailbox.push(format!("message-{index}")).unwrap(); } + assert!(mailbox.push("overflow".into()).is_err()); let mut receiver = mailbox.try_acquire().unwrap(); - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { assert_eq!( receiver.recv().await, Some(format!("message-{index}")), @@ -641,17 +891,21 @@ mod tests { async fn outbound_mailbox_does_not_stall_when_subscriber_is_slow() { let mailbox = OutboundMailbox::new(); let mut receiver = mailbox.try_acquire().unwrap(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { mailbox.push(format!("message-{index}")).unwrap(); } - for index in 0..ISSUE_288_BURST { + assert!(mailbox.push("overflow".into()).is_err()); + for index in 0..capacity { assert_eq!( receiver.recv().await, Some(format!("message-{index}")), "message {index} should remain ordered" ); } + mailbox.push("recovered".into()).unwrap(); + assert_eq!(receiver.recv().await.as_deref(), Some("recovered")); } #[tokio::test] @@ -679,6 +933,7 @@ mod tests { #[tokio::test] async fn slow_session_mailbox_does_not_stall_other_routes() { let outbound = HttpOutbound::new(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; let mut slow_session = outbound .session_stream("slow-session") .await @@ -691,7 +946,7 @@ mod tests { .unwrap(); timeout(Duration::from_secs(1), async { - for index in 0..ISSUE_288_BURST { + for index in 0..=capacity { let message = RawJsonRpcMessage::notification( "session/update".to_string(), serde_json::json!({ @@ -701,7 +956,15 @@ mod tests { ) .unwrap(); let serialized = serde_json::to_string(&message).unwrap(); - outbound.route_outbound(&message, serialized).await.unwrap(); + let result = outbound.route_outbound(&message, serialized).await; + if index == capacity { + assert!( + result.is_err(), + "overflow must be explicit, not silently dropped" + ); + } else { + result.unwrap(); + } } let marker = RawJsonRpcMessage::notification( @@ -727,7 +990,7 @@ mod tests { true ); - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { let message = slow_session.recv().await.unwrap(); assert_eq!( serde_json::from_str::(&message).unwrap()["params"]["index"], @@ -772,10 +1035,11 @@ mod tests { let future = Box::pin(async move { agent .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( RequestId::Number(1), Ok(serde_json::json!({ "done": true })), ))) + .await .unwrap(); Ok(()) }); @@ -801,11 +1065,12 @@ mod tests { emit.notified().await; agent .tx - .unbounded_send(TransportFrame::Malformed { + .send_frame(TransportFrame::Malformed { raw: "{not json".to_string(), error: agent_client_protocol::Error::parse_error() .data("transport parse error"), }) + .await .unwrap(); std::future::pending::>().await }); @@ -832,7 +1097,8 @@ mod tests { let future = Box::pin(async move { agent .tx - .unbounded_send(TransportFrame::Single(message)) + .send_frame(TransportFrame::Single(message)) + .await .unwrap(); exit.notified().await; Ok(()) @@ -871,7 +1137,8 @@ mod tests { .expect("test batch is non-empty"); agent .tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .unwrap(); exit.notified().await; Ok(()) @@ -898,13 +1165,14 @@ mod tests { emit.notified().await; agent .tx - .unbounded_send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::notification( "test/final".to_string(), serde_json::json!({}), ) .unwrap(), )) + .await .unwrap(); Ok(()) }); @@ -974,13 +1242,11 @@ mod tests { .expect("buffered response should be forwarded before teardown"); assert!(matches!( - frame, - TransportFrame::Single(RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { - id: RequestId::Number(1), - .. - } - )) + frame.frame(), + TransportFrame::Single(RawJsonRpcMessage::Response(RpcResponse::Result { + id: RequestId::Number(1), + .. + })) )); timeout(Duration::from_secs(1), async { loop { @@ -1028,6 +1294,87 @@ mod tests { )); } + #[tokio::test] + async fn fatal_router_overflow_tears_down_idle_agent_and_metadata() { + struct BurstThenIdle(Arc); + struct Dropped(Arc); + impl Drop for Dropped { + fn drop(&mut self) { + self.0.store(true, std::sync::atomic::Ordering::SeqCst); + } + } + impl AgentFactory for BurstThenIdle { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (agent, transport) = Channel::duplex(); + let dropped = Dropped(self.0.clone()); + let future = Box::pin(async move { + let _dropped = dropped; + for _ in 0..33 { + agent + .tx + .send_frame(TransportFrame::Single( + RawJsonRpcMessage::notification( + "test/burst".into(), + serde_json::json!({}), + ) + .unwrap(), + )) + .await + .unwrap(); + } + std::future::pending::<()>().await; + Ok(()) + }); + (transport, future) + } + } + + let dropped = Arc::new(std::sync::atomic::AtomicBool::new(false)); + let registry = ConnectionRegistry::new(Arc::new(BurstThenIdle(dropped.clone()))); + let (id, connection) = registry.create_connection().await; + let source = connection + .admit_frame_to_agent(TransportFrame::Single( + RawJsonRpcMessage::notification("test/source".into(), serde_json::json!({})) + .unwrap(), + )) + .unwrap(); + connection + .register_post_routes( + &["S".into()], + &[(RequestId::Number(1), ResponseRoute::Session("S".into()))], + source.permit(), + ) + .await + .unwrap(); + let OutboundTransport::Http(http) = &connection.outbound_transport else { + unreachable!() + }; + assert_eq!(http.pending_routes.lock().await.len(), 1); + drop(source); + connection.start_router().await; + timeout(Duration::from_secs(1), async { + let mut closed = connection.subscribe_closed(); + while !*closed.borrow() { + closed.changed().await.unwrap(); + } + while registry.get(&id).await.is_some() + || !dropped.load(std::sync::atomic::Ordering::SeqCst) + || !http.pending_routes.lock().await.is_empty() + || connection.subscribe_session_stream("S").await.is_some() + { + sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("fatal router exit must stop idle agent and remove registry entry"); + assert!(http.pending_routes.lock().await.is_empty()); + } + #[tokio::test] async fn protocol_level_notification_routes_to_connection_stream() { let exit = Arc::new(Notify::new()); @@ -1045,6 +1392,7 @@ mod tests { })); let (_connection_id, connection) = registry.create_connection().await; let mut connection_rx = connection.subscribe_connection_stream().unwrap(); + connection.ensure_session("session-1").await; let mut session_rx = connection .subscribe_session_stream("session-1") .await @@ -1118,7 +1466,7 @@ mod tests { let serialized = serde_json::to_string(&batch).unwrap(); outbound - .route_outbound_batch(&batch, serialized.clone()) + .route_outbound_batch(&batch, serialized.clone(), None) .await .unwrap(); diff --git a/src/agent-client-protocol-http/src/connection_admission_tests.rs b/src/agent-client-protocol-http/src/connection_admission_tests.rs new file mode 100644 index 00000000..c0f6838f --- /dev/null +++ b/src/agent-client-protocol-http/src/connection_admission_tests.rs @@ -0,0 +1,232 @@ +use agent_client_protocol::ConnectionLimits; +use serde_json::json; + +use super::*; + +#[tokio::test] +async fn outbound_lease_releases_delivered_frame_before_waiting_or_idle_poll() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".into(), json!({"data": "x".repeat(100)})) + .unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes + 1, + max_queued_bytes: bytes * 2 + 1, + max_queued_frames: 2, + }; + let (_, channel) = Channel::duplex_with_limits(limits); + let admission = channel.tx.admission(); + let mailbox = OutboundMailbox::new(); + let mut lease = mailbox.try_acquire().unwrap(); + let (_, first) = admission.try_admit(frame.clone()).unwrap().into_parts(); + mailbox + .push_with_permit("first".into(), Some(first)) + .unwrap(); + assert_eq!(lease.recv().await.as_deref(), Some("first")); + assert!(admission.try_admit(frame.clone()).is_err()); + assert!(lease.try_recv().is_err()); + let (_, second) = admission.try_admit(frame.clone()).unwrap().into_parts(); + mailbox + .push_with_permit("second".into(), Some(second)) + .unwrap(); + assert_eq!(lease.try_recv().unwrap(), "second"); + assert!(admission.try_admit(frame.clone()).is_err()); + let wait = tokio::spawn(async move { lease.recv().await }); + tokio::task::yield_now().await; + let (_, third) = admission.try_admit(frame).unwrap().into_parts(); + mailbox + .push_with_permit("third".into(), Some(third)) + .unwrap(); + assert_eq!(wait.await.unwrap().as_deref(), Some("third")); +} + +#[tokio::test] +async fn session_key_does_not_pin_unrelated_post_payload() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification( + "session/update".into(), + json!({"sessionId": "persisted", "payload": "x".repeat(512)}), + ) + .unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2 + 64, + max_queued_frames: 2, + }; + let (_, channel) = Channel::duplex_with_limits(limits); + let admission = channel.tx.admission(); + let (_, source) = admission.try_admit(frame.clone()).unwrap().into_parts(); + let mut http = HttpOutbound::new(); + http.limits = limits; + http.register_post_routes(&["persisted".into()], &[], &source) + .await + .unwrap(); + drop(source); + assert!(http.session_streams.read().await.contains_key("persisted")); + assert!( + admission.try_admit(frame).is_ok(), + "the retained session key must not pin its source payload" + ); +} + +#[tokio::test] +async fn failed_session_key_reservation_does_not_publish_partial_batch() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/data".into(), json!({"payload": "x".repeat(512)})) + .unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2 + 12, + max_queued_frames: 3, + }; + let (_, channel) = Channel::duplex_with_limits(limits); + let (_, source) = channel + .tx + .admission() + .try_admit(frame) + .unwrap() + .into_parts(); + let mut http = HttpOutbound::new(); + http.limits = limits; + assert!( + http.register_post_routes(&["one".into(), "another".into()], &[], &source) + .await + .is_err() + ); + assert!(http.session_streams.read().await.is_empty()); + assert!(http.pending_routes.lock().await.is_empty()); +} + +#[tokio::test] +async fn route_and_session_metadata_admission_is_atomic_and_releases_permits() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::request("test/request".into(), json!({}), RequestId::Number(1)).unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes + 16, + max_queued_bytes: bytes * 2 + 32, + max_queued_frames: 2, + }; + let (_caller, transport) = Channel::duplex_with_limits(limits); + let admission = transport.tx.admission(); + let (_, permit) = admission.try_admit(frame.clone()).unwrap().into_parts(); + let mut http = HttpOutbound::new(); + http.limits = limits; + assert!( + OutboundTransport::Http(Box::new(HttpOutbound::new())) + .subscribe_session_stream("unknown") + .await + .is_none() + ); + let first = [(RequestId::Number(1), ResponseRoute::Session("one".into()))]; + http.register_post_routes(&["one".into()], &first, &permit) + .await + .unwrap(); + let extra = [(RequestId::Number(2), ResponseRoute::Session("two".into()))]; + assert!( + http.register_post_routes(&["two".into()], &extra, &permit) + .await + .is_err() + ); + assert_eq!(http.session_streams.read().await.len(), 1); + assert_eq!(http.pending_routes.lock().await.len(), 1); + drop(permit); + assert!(admission.try_admit(frame.clone()).is_err()); + assert_eq!( + take_pending_route( + &mut *http.pending_routes.lock().await, + &RequestId::Number(1) + ), + Some(ResponseRoute::Session("one".into())) + ); + assert!( + admission.try_admit(frame.clone()).is_ok(), + "removing a pending route releases its full source-frame charge" + ); + http.session_streams.write().await.clear(); + assert!(admission.try_admit(frame).is_ok()); +} + +#[tokio::test] +async fn concurrent_post_reservations_preserve_adopted_session() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::request("test/request".into(), json!({}), RequestId::Number(1)).unwrap(), + ); + let (_caller, transport) = Channel::duplex(); + let (inbound_tx, mut inbound_rx) = mpsc::channel(2); + let connection = Connection { + inbound_tx, + inbound_admission: transport.tx.admission(), + outbound_rx: Mutex::new(None), + agent_handle: Mutex::new(None), + router_handle: Mutex::new(None), + closed_tx: watch::channel(false).0, + outbound_transport: OutboundTransport::http(), + }; + let a = connection.reserve_inbound().unwrap(); + let b = connection.reserve_inbound().unwrap(); + assert!(connection.reserve_inbound().is_err()); + + // A publishes S first; B adopts S and enqueues while A is paused. + // Neither commit may subsequently fail queue admission or erase S. + let a_frame = connection.admit_frame_to_agent(frame.clone()).unwrap(); + connection + .register_post_routes(&["S".into()], &[], a_frame.permit()) + .await + .unwrap(); + let OutboundTransport::Http(http) = &connection.outbound_transport else { + unreachable!("HTTP test connection"); + }; + let original = http.session_streams.read().await["S"].0.clone(); + let b_frame = connection.admit_frame_to_agent(frame).unwrap(); + connection + .register_post_routes(&["S".into()], &[], b_frame.permit()) + .await + .unwrap(); + assert!(Arc::ptr_eq( + &original, + &http.session_streams.read().await["S"].0 + )); + b.send(b_frame); + a.send(a_frame); + assert!(inbound_rx.recv().await.is_some()); + assert!(inbound_rx.recv().await.is_some()); + assert!(connection.subscribe_session_stream("S").await.is_some()); + + // A cancelled before metadata registration cannot strand a queue slot. + let cancelled = connection.reserve_inbound().unwrap(); + drop(cancelled); + assert!(connection.reserve_inbound().is_ok()); +} + +#[tokio::test] +async fn successful_session_response_registers_stream_before_get() { + let response = RawJsonRpcMessage::response( + RequestId::Number(1), + Ok(json!({"sessionId": "new-session"})), + ); + let frame = TransportFrame::Single(response.clone()); + let (_, channel) = Channel::duplex(); + let (_, permit) = channel + .tx + .admission() + .try_admit(frame.clone()) + .unwrap() + .into_parts(); + let http = HttpOutbound::new(); + http.route_outbound_with_permit(&response, frame.to_json().unwrap(), Some(permit)) + .await + .unwrap(); + assert!( + http.session_streams + .read() + .await + .contains_key("new-session") + ); +} diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index 525b8e7c..016a2937 100644 --- a/src/agent-client-protocol-http/src/http_server.rs +++ b/src/agent-client-protocol-http/src/http_server.rs @@ -1,8 +1,8 @@ use std::{convert::Infallible, error::Error as _, sync::Arc, time::Duration}; use agent_client_protocol::{ - RawJsonRpcMessage, TransportBatchEntry, TransportFrame, schema::v1::RequestId, - schema::v1::Response as RpcResponse, + RawJsonRpcMessage, RawJsonRpcResponse as RpcResponse, TransportBatchEntry, TransportFrame, + schema::v1::RequestId, }; use axum::{ body::Body, @@ -102,7 +102,9 @@ pub(crate) async fn handle_post( ) .into_response(); }; - if let Some(initialize_failed) = initialize_response_failed(&frame, &initialize_id) { + if let Some(initialize_failed) = + initialize_response_failed(frame.frame(), &initialize_id) + { break (frame, initialize_failed); } @@ -114,7 +116,7 @@ pub(crate) async fn handle_post( return (StatusCode::INTERNAL_SERVER_ERROR, error).into_response(); } }; - let init_response = match init_response_frame.to_json() { + let init_response = match init_response_frame.frame().to_json() { Ok(response) => response, Err(e) => { initialize_cleanup.cleanup().await; @@ -170,16 +172,26 @@ pub(crate) async fn handle_post( } } - for session_id in session_routes { - connection.ensure_session(&session_id).await; - } - for (request_id, route) in pending_routes { - connection.record_pending_route(request_id, route).await; - } - - if connection.send_frame_to_agent(frame).is_err() { - return StatusCode::INTERNAL_SERVER_ERROR.into_response(); + let admitted = match connection.admit_frame_to_agent(frame) { + Ok(frame) => frame, + Err(error) => return (StatusCode::TOO_MANY_REQUESTS, error).into_response(), + }; + // Claim the queue slot before publishing session and response-route + // metadata. After publication, sending through this permit cannot fail + // due to another POST filling the queue. + let inbound_slot = match connection.reserve_inbound() { + Ok(slot) => slot, + Err(error) => return (StatusCode::TOO_MANY_REQUESTS, error).into_response(), + }; + let permit = admitted.permit().clone(); + if let Err(error) = connection + .register_post_routes(&session_routes, &pending_routes, &permit) + .await + { + return (StatusCode::TOO_MANY_REQUESTS, error).into_response(); } + drop(permit); + inbound_slot.send(admitted); StatusCode::ACCEPTED.into_response() } @@ -356,7 +368,7 @@ pub(crate) async fn handle_get( let Some(mut receiver) = receiver else { return ( StatusCode::CONFLICT, - "outbound stream already has a subscriber", + "outbound stream missing or already has a subscriber", ) .into_response(); }; @@ -480,8 +492,8 @@ mod tests { use std::sync::Arc; use agent_client_protocol::{ - Channel, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame, - schema::v1::RequestId, + BudgetedFrame, Channel, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, + TransportFrame, schema::v1::RequestId, }; use futures::{StreamExt, future::BoxFuture}; use serde_json::json; @@ -493,8 +505,6 @@ mod tests { use super::*; use crate::connection::AgentFactory; - const ISSUE_288_BURST: usize = 1_025; - struct CapturingAgentFactory { forwarded: mpsc::UnboundedSender, } @@ -514,7 +524,7 @@ mod tests { tx: _, } = agent; while let Some(frame) = incoming.next().await { - let TransportFrame::Single(message) = frame else { + let TransportFrame::Single(message) = frame.into_frame() else { panic!("expected a single JSON-RPC frame"); }; if forwarded.send(message).is_err() { @@ -528,6 +538,54 @@ mod tests { } } + #[tokio::test] + async fn rejected_post_does_not_remove_accepted_session_stream() { + let (forwarded, _receiver) = mpsc::unbounded_channel(); + let registry = Arc::new(ConnectionRegistry::new(Arc::new(CapturingAgentFactory { + forwarded, + }))); + let (id, connection) = registry.create_connection().await; + let post = || { + Request::builder() + .method("POST") + .uri("/acp") + .header(header::CONTENT_TYPE, JSON_MIME_TYPE) + .header(HEADER_CONNECTION_ID, id.as_str()) + .header(HEADER_SESSION_ID, "S") + .body(Body::from( + json!({"jsonrpc":"2.0","method":"session/update","params":{}}).to_string(), + )) + .unwrap() + }; + assert_eq!( + handle_post(State(registry.clone()), post()).await.status(), + StatusCode::ACCEPTED + ); + let mut slots = Vec::new(); + while let Ok(slot) = connection.reserve_inbound() { + slots.push(slot); + } + assert!(!slots.is_empty()); + assert_eq!( + handle_post(State(registry.clone()), post()).await.status(), + StatusCode::TOO_MANY_REQUESTS + ); + let get = Request::builder() + .uri("/acp") + .header(header::ACCEPT, EVENT_STREAM_MIME_TYPE) + .header(HEADER_CONNECTION_ID, id.as_str()) + .header(HEADER_SESSION_ID, "S") + .body(Body::empty()) + .unwrap(); + assert_eq!( + handle_get(registry.clone(), get).await.status(), + StatusCode::OK + ); + drop(slots); + registry.remove(&id).await; + connection.shutdown().await; + } + struct RejectingInitializeAgentFactory; impl AgentFactory for RejectingInitializeAgentFactory { @@ -539,15 +597,16 @@ mod tests { ) { let (mut agent, transport) = Channel::duplex(); let future = Box::pin(async move { - match agent.rx.next().await { + match agent.rx.next().await.map(BudgetedFrame::into_frame) { Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) => { agent .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Err(agent_client_protocol::Error::invalid_request() .data("initialize rejected")), ))) + .await .unwrap(); } Some(TransportFrame::Batch(batch)) => { @@ -567,10 +626,11 @@ mod tests { }); agent .tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages(responses) .expect("request batch has responses"), )) + .await .unwrap(); } Some(TransportFrame::Single(_) | TransportFrame::Malformed { .. }) | None => {} @@ -619,7 +679,9 @@ mod tests { let (mut agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); let future = Box::pin(async move { - let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { + let Some(TransportFrame::Batch(batch)) = + agent.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected one batch frame"); }; let mut methods = Vec::new(); @@ -666,7 +728,8 @@ mod tests { TransportBatch::from_messages(responses).expect("responses are non-empty"); agent .tx - .unbounded_send(TransportFrame::Batch(responses)) + .send_frame(TransportFrame::Batch(responses)) + .await .unwrap(); std::future::pending::>().await }); @@ -686,7 +749,9 @@ mod tests { ) { let (mut agent, transport) = Channel::duplex(); let future = Box::pin(async move { - let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { + let Some(TransportFrame::Batch(batch)) = + agent.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected one initial batch frame"); }; let responses = batch.entries().filter_map(|entry| { @@ -702,20 +767,22 @@ mod tests { agent .tx - .unbounded_send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::notification( "custom/during-initialize".into(), json!({ "phase": "before-response" }), ) .expect("test notification should serialize"), )) + .await .unwrap(); agent .tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages(responses) .expect("initial batch has response-bearing requests"), )) + .await .unwrap(); std::future::pending::>().await }); @@ -989,6 +1056,7 @@ mod tests { }))); let (connection_id, connection) = registry.create_connection().await; let mut connection_outbound = connection.subscribe_connection_stream().unwrap(); + connection.ensure_session("session-1").await; let mut session_outbound = connection .subscribe_session_stream("session-1") .await @@ -1057,6 +1125,7 @@ mod tests { }))); let (connection_id, connection) = registry.create_connection().await; let mut connection_outbound = connection.subscribe_connection_stream().unwrap(); + connection.ensure_session("session-1").await; let mut session_outbound = connection .subscribe_session_stream("session-1") .await @@ -1258,8 +1327,9 @@ mod tests { } #[tokio::test] - async fn sse_buffers_burst_without_polling_slow_subscriber() { + async fn sse_bounds_burst_and_drains_every_accepted_message() { let (forwarded_tx, _forwarded_rx) = mpsc::unbounded_channel(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; let registry = Arc::new(ConnectionRegistry::new(Arc::new(CapturingAgentFactory { forwarded: forwarded_tx, }))); @@ -1275,11 +1345,16 @@ mod tests { assert_eq!(response.status(), StatusCode::OK); timeout(Duration::from_secs(1), async { - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { connection .push_connection_stream_for_test(format!("message-{index}")) .unwrap(); } + assert!( + connection + .push_connection_stream_for_test("overflow".into()) + .is_err() + ); }) .await .expect("enqueueing must not wait for the SSE body to be polled"); @@ -1297,7 +1372,7 @@ mod tests { .lines() .filter_map(|line| line.strip_prefix("data: ")) .collect::>(); - let expected = (0..ISSUE_288_BURST) + let expected = (0..capacity) .map(|index| format!("message-{index}")) .collect::>(); assert_eq!( @@ -1313,6 +1388,8 @@ mod tests { forwarded: forwarded_tx, }))); let (connection_id, connection) = registry.create_connection().await; + connection.ensure_session("session-1").await; + connection.ensure_session("session-2").await; let request = |session_id: Option<&str>| { let mut request = Request::builder() .method("GET") diff --git a/src/agent-client-protocol-http/src/protocol.rs b/src/agent-client-protocol-http/src/protocol.rs index 79407adf..5e2862f3 100644 --- a/src/agent-client-protocol-http/src/protocol.rs +++ b/src/agent-client-protocol-http/src/protocol.rs @@ -47,7 +47,7 @@ pub(crate) fn is_connection_scoped_protocol_message(msg: &RawJsonRpcMessage) -> || is_cancel_request_message(msg) } -fn is_cancel_request_message(msg: &RawJsonRpcMessage) -> bool { +pub(crate) fn is_cancel_request_message(msg: &RawJsonRpcMessage) -> bool { let RawJsonRpcMessage::Notification(notification) = msg else { return false; }; diff --git a/src/agent-client-protocol-http/src/server.rs b/src/agent-client-protocol-http/src/server.rs index 695b1699..07cab781 100644 --- a/src/agent-client-protocol-http/src/server.rs +++ b/src/agent-client-protocol-http/src/server.rs @@ -196,9 +196,190 @@ async fn handle_get( #[cfg(test)] mod tests { use super::*; + use agent_client_protocol::{ + Channel, ConnectTo, RawJsonRpcMessage, TransportBatch, TransportFrame, + schema::v1::RequestId, + }; use axum::body::Body; + use futures::{StreamExt, future::BoxFuture}; + use serde_json::json; + use tokio::{ + net::TcpListener, + time::{Duration, timeout}, + }; use tower::{Layer as _, ServiceExt as _, service_fn}; + struct HistoryAgent; + + impl crate::connection::AgentFactory for HistoryAgent { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (mut agent, transport) = Channel::duplex(); + let run = Box::pin(async move { + while let Some(frame) = agent.rx.next().await { + let messages = match frame.into_frame() { + TransportFrame::Single(message) => vec![message], + TransportFrame::Batch(batch) => batch + .entries() + .filter_map(|entry| match entry { + agent_client_protocol::TransportBatchEntry::Message(message) => { + Some(message.clone()) + } + agent_client_protocol::TransportBatchEntry::Malformed { + .. + } => None, + }) + .collect(), + TransportFrame::Malformed { .. } => continue, + }; + for message in messages { + let RawJsonRpcMessage::Request(request) = message else { + continue; + }; + if request.method.as_ref() != "initialize" { + let Some(agent_client_protocol::RawJsonRpcParams::Object(params)) = + request.params.as_ref() + else { + panic!("session request must have object params"); + }; + for index in 0..2 { + agent + .tx + .send_frame(TransportFrame::Single( + RawJsonRpcMessage::notification( + "session/update".into(), + json!({"sessionId": params["sessionId"], "index": index}), + ) + .unwrap(), + )) + .await + .unwrap(); + } + } + agent + .tx + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( + request.id, + Ok(json!({})), + ))) + .await + .unwrap(); + } + } + Ok(()) + }); + (transport, run) + } + } + + #[tokio::test] + async fn cold_session_post_registers_stream_before_history_for_single_and_batch() { + let registry = Arc::new(ConnectionRegistry::new(Arc::new(HistoryAgent))); + let app = AcpHttpServer { + registry: registry.clone(), + options: ServerOptions::default(), + } + .into_router(); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + for (index, method) in ["session/load", "session/resume"].into_iter().enumerate() { + let client = crate::client::HttpClient::new(format!("http://{address}")).unwrap(); + let (mut caller, driver) = client.into_channel_and_future(); + let driver = tokio::spawn(driver); + caller + .tx + .send_frame(TransportFrame::Single( + RawJsonRpcMessage::request( + "initialize".into(), + json!({}), + RequestId::Number(1), + ) + .unwrap(), + )) + .await + .unwrap(); + let init = timeout(Duration::from_secs(2), caller.rx.next()) + .await + .unwrap() + .unwrap(); + assert!(matches!( + init.frame(), + TransportFrame::Single(RawJsonRpcMessage::Response(_)) + )); + drop(init); + let request = RawJsonRpcMessage::request( + method.into(), + json!({"sessionId": "persisted"}), + RequestId::Number(2), + ) + .unwrap(); + let frame = if index == 0 { + TransportFrame::Single(request) + } else { + let second = RawJsonRpcMessage::request( + method.into(), + json!({"sessionId": "other-persisted"}), + RequestId::Number(3), + ) + .unwrap(); + TransportFrame::Batch(TransportBatch::from_messages([request, second]).unwrap()) + }; + caller.tx.send_frame(frame).await.unwrap(); + let mut seen = std::collections::BTreeMap::>::new(); + let session_count = index + 1; + let mut responses = 0; + for _ in 0..session_count * 3 { + let frame = timeout(Duration::from_secs(3), caller.rx.next()) + .await + .unwrap() + .unwrap(); + match frame.frame() { + TransportFrame::Single(RawJsonRpcMessage::Notification(notification)) => { + let agent_client_protocol::RawJsonRpcParams::Object(params) = + notification.params.as_ref().unwrap() + else { + panic!("history update must have object params"); + }; + seen.entry(params["sessionId"].as_str().unwrap().to_owned()) + .or_default() + .push(params["index"].as_u64().unwrap()); + } + TransportFrame::Single(RawJsonRpcMessage::Response(_)) => { + let response: serde_json::Value = + serde_json::from_str(&frame.frame().to_json().unwrap()).unwrap(); + let session = match response["id"].as_u64().unwrap() { + 2 => "persisted", + 3 => "other-persisted", + id => panic!("unexpected response ID: {id}"), + }; + assert_eq!( + seen.get(session).map(Vec::as_slice), + Some([0, 1].as_slice()), + "{method} must deliver history before its response" + ); + responses += 1; + } + other => panic!("unexpected history frame: {other:?}"), + } + } + assert_eq!(seen.len(), session_count); + assert_eq!(responses, session_count); + drop(caller); + timeout(Duration::from_secs(3), driver) + .await + .unwrap() + .unwrap() + .unwrap(); + } + assert_eq!(registry.len().await, 0); + server.abort(); + } + #[test] fn cors_is_disabled_by_default() { assert_eq!(ServerOptions::default().cors, CorsOptions::Disabled); diff --git a/src/agent-client-protocol-http/src/websocket_server.rs b/src/agent-client-protocol-http/src/websocket_server.rs index d051f4d0..08356125 100644 --- a/src/agent-client-protocol-http/src/websocket_server.rs +++ b/src/agent-client-protocol-http/src/websocket_server.rs @@ -162,13 +162,23 @@ where { trace!(connection_id = %connection_id, session_id = %sid, request_id = ?req.id, "Client → Agent (session)"); } - if connection.send_frame_to_agent(frame).is_err() { - error!(connection_id = %connection_id, "Agent channel closed"); - drain_outbound_until_closed(ws_tx, outbound_rx, closed, connection_id).await; - false - } else { - true + let frame = match connection.admit_frame_to_agent(frame) { + Ok(frame) => frame, + Err(error) => { + warn!(connection_id = %connection_id, "Rejecting WebSocket frame: {error}"); + return false; + } + }; + if let Err(error) = connection.send_budgeted_frame_to_agent(frame) { + if connection.agent_channel_closed() { + error!(connection_id = %connection_id, "Agent channel closed"); + drain_outbound_until_closed(ws_tx, outbound_rx, closed, connection_id).await; + } else { + warn!(connection_id = %connection_id, "Rejecting WebSocket frame: {error}"); + } + return false; } + true } async fn drain_outbound_until_closed( @@ -235,8 +245,8 @@ where #[cfg(test)] mod tests { use agent_client_protocol::{ - Channel, TransportBatch, TransportBatchEntry, TransportFrame, - schema::v1::{RequestId, Response as RpcResponse}, + BudgetedFrame, Channel, RawJsonRpcResponse as RpcResponse, TransportBatch, + TransportBatchEntry, TransportFrame, schema::v1::RequestId, }; use async_tungstenite::{tokio::connect_async, tungstenite::Message as ClientWsMessage}; use axum::{Router, extract::WebSocketUpgrade, routing::get}; @@ -252,8 +262,6 @@ mod tests { use super::*; - const ISSUE_288_BURST: usize = 1_025; - struct CapturingAgentFactory { forwarded: mpsc::UnboundedSender, } @@ -273,7 +281,7 @@ mod tests { tx: outgoing, } = agent; while let Some(frame) = incoming.next().await { - match frame { + match frame.into_frame() { TransportFrame::Single(message) => { if forwarded.send(message).is_err() { break; @@ -281,9 +289,11 @@ mod tests { } TransportFrame::Malformed { error, .. } => { outgoing - .unbounded_send(TransportFrame::Single( - RawJsonRpcMessage::response(RequestId::Null, Err(error)), - )) + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( + RequestId::Null, + Err(error), + ))) + .await .unwrap(); } TransportFrame::Batch(_) => panic!("expected a single JSON-RPC frame"), @@ -296,6 +306,161 @@ mod tests { } } + struct LimitedAgentFactory { + forwarded: mpsc::UnboundedSender, + } + + struct StalledAgentFactory; + + impl AgentFactory for StalledAgentFactory { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (agent, transport) = + Channel::duplex_with_limits(agent_client_protocol::ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 2048, + max_queued_frames: 4, + }); + let future = Box::pin(async move { + std::future::pending::<()>().await; + drop(agent); + Ok(()) + }); + (transport, future) + } + } + + impl AgentFactory for LimitedAgentFactory { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let limits = agent_client_protocol::ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 2048, + max_queued_frames: 4, + }; + let (mut agent, transport) = Channel::duplex_with_limits(limits); + let forwarded = self.forwarded.clone(); + let future = Box::pin(async move { + while let Some(frame) = agent.rx.next().await { + if let TransportFrame::Single(message) = frame.into_frame() { + forwarded.send(message).ok(); + } + } + Ok(()) + }); + (transport, future) + } + } + + #[tokio::test] + async fn oversized_websocket_frame_closes_live_agent_connection_and_registry() { + let (forwarded_tx, mut forwarded_rx) = mpsc::unbounded_channel(); + let registry = Arc::new(ConnectionRegistry::new(Arc::new(LimitedAgentFactory { + forwarded: forwarded_tx, + }))); + let app = Router::new().route( + "/acp", + get({ + let registry = registry.clone(); + move |ws: WebSocketUpgrade| { + let registry = registry.clone(); + async move { handle_ws_upgrade(registry, ws) } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let (mut socket, _) = connect_async(format!("ws://{addr}/acp")).await.unwrap(); + let ordinary = serde_json::json!({ + "jsonrpc": "2.0", "method": "test/valid", "params": {} + }); + socket + .send(ClientWsMessage::Text(ordinary.to_string().into())) + .await + .unwrap(); + timeout(Duration::from_secs(2), forwarded_rx.recv()) + .await + .unwrap() + .expect("live agent receives first message"); + let oversized = serde_json::json!({ + "jsonrpc": "2.0", "method": "test/oversized", + "params": { "payload": "x".repeat(512) } + }); + socket + .send(ClientWsMessage::Text(oversized.to_string().into())) + .await + .unwrap(); + timeout(Duration::from_secs(2), async { + loop { + if registry.len().await == 0 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("rejected frame must release registry entry"); + assert!( + timeout(Duration::from_secs(2), socket.next()).await.is_ok(), + "rejected frame must terminate the socket" + ); + server.abort(); + } + + #[tokio::test] + async fn saturated_websocket_frame_closes_stalled_agent_connection() { + let registry = Arc::new(ConnectionRegistry::new(Arc::new(StalledAgentFactory))); + let app = Router::new().route( + "/acp", + get({ + let registry = registry.clone(); + move |ws: WebSocketUpgrade| { + let registry = registry.clone(); + async move { handle_ws_upgrade(registry, ws) } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() }); + let (mut socket, _) = connect_async(format!("ws://{addr}/acp")).await.unwrap(); + let frame = json!({ + "jsonrpc": "2.0", "method": "test/stalled", "params": {"payload": "x".repeat(400)} + }) + .to_string(); + for _ in 0..4 { + // The agent deliberately does not consume its input. + if socket + .send(ClientWsMessage::Text(frame.clone().into())) + .await + .is_err() + { + break; + } + } + timeout(Duration::from_secs(2), async { + loop { + if registry.len().await == 0 { + break; + } + tokio::time::sleep(Duration::from_millis(10)).await; + } + }) + .await + .expect("saturated admission must release registry entry"); + assert!(timeout(Duration::from_secs(2), socket.next()).await.is_ok()); + server.abort(); + } + struct BatchAgentFactory { forwarded: mpsc::UnboundedSender>, } @@ -310,7 +475,9 @@ mod tests { let (mut agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); let future = Box::pin(async move { - let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { + let Some(TransportFrame::Batch(batch)) = + agent.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected one batch frame"); }; let mut methods = Vec::new(); @@ -331,7 +498,8 @@ mod tests { TransportBatch::from_messages(responses).expect("responses are non-empty"); agent .tx - .unbounded_send(TransportFrame::Batch(responses)) + .send_frame(TransportFrame::Batch(responses)) + .await .unwrap(); std::future::pending::>().await }); @@ -357,13 +525,14 @@ mod tests { emit.notified().await; agent .tx - .unbounded_send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::notification( "test/final".to_string(), serde_json::json!({}), ) .unwrap(), )) + .await .unwrap(); Ok(()) }); @@ -390,13 +559,14 @@ mod tests { emit.notified().await; agent .tx - .unbounded_send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::notification( "test/final".to_string(), serde_json::json!({}), ) .unwrap(), )) + .await .unwrap(); Ok(()) }); @@ -406,8 +576,9 @@ mod tests { } #[tokio::test] - async fn websocket_buffers_burst_without_polling_slow_subscriber() { + async fn websocket_bounds_burst_and_drains_every_accepted_message() { let (forwarded_tx, _forwarded_rx) = mpsc::unbounded_channel(); + let capacity = agent_client_protocol::ConnectionLimits::default().max_queued_frames; let registry = Arc::new(ConnectionRegistry::new(Arc::new(CapturingAgentFactory { forwarded: forwarded_tx, }))); @@ -425,11 +596,16 @@ mod tests { .await; let mut outbound_rx = connection.subscribe_all_outbound().unwrap(); - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { connection .push_all_outbound_for_test(format!("message-{index}")) .unwrap(); } + assert!( + connection + .push_all_outbound_for_test("overflow".into()) + .is_err() + ); let mut closed = connection.subscribe_closed(); let (mut ws_tx, mut ws_rx) = socket.split(); @@ -457,7 +633,7 @@ mod tests { let (mut client, _) = connect_async(format!("ws://{addr}/acp")).await.unwrap(); timeout(Duration::from_secs(5), async { - for index in 0..ISSUE_288_BURST { + for index in 0..capacity { let frame = client.next().await.unwrap().unwrap(); let ClientWsMessage::Text(text) = frame else { panic!("expected text frame: {frame:?}"); @@ -466,7 +642,7 @@ mod tests { } }) .await - .expect("WebSocket should deliver the complete burst"); + .expect("WebSocket should deliver every accepted frame"); server.abort(); } diff --git a/src/agent-client-protocol-polyfill/CHANGELOG.md b/src/agent-client-protocol-polyfill/CHANGELOG.md index 77cb9487..c02f4ee7 100644 --- a/src/agent-client-protocol-polyfill/CHANGELOG.md +++ b/src/agent-client-protocol-polyfill/CHANGELOG.md @@ -7,6 +7,21 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ## [Unreleased] +### Changed + +- Replace the MCP session bridge with a latest-only MCP 2026-07-28 HTTP adapter: + one native ACP request per POST, request-scoped SSE, and stream-close + cancellation. Remove the initialize/connect/disconnect, GET, batch, and + session-header paths; existing older MCP HTTP clients must be upgraded. +- Require runtime bearer credentials from the rewritten declaration, validate + Origin and mirrored routing headers, and preserve logical request identities + independently of overlapping HTTP IDs. +- Bound notification queues and request/listener admission. Translate + subscription IDs in namespaced MCP metadata for notifications and completion. +- Fail closed for unsupported `x-mcp-header` tools. Direct tool calls perform + descriptor lookup internally rather than requiring a client-side list + handshake. This remains a draft, not a claim of full HTTP conformance. + ## [2.2.0](https://github.com/agentclientprotocol/rust-sdk/compare/agent-client-protocol-polyfill-v2.1.0...agent-client-protocol-polyfill-v2.2.0) - 2026-09-18 ### Other diff --git a/src/agent-client-protocol-polyfill/Cargo.toml b/src/agent-client-protocol-polyfill/Cargo.toml index 6c678b4f..cba60257 100644 --- a/src/agent-client-protocol-polyfill/Cargo.toml +++ b/src/agent-client-protocol-polyfill/Cargo.toml @@ -19,14 +19,17 @@ unstable_session_fork = ["agent-client-protocol/unstable_session_fork"] agent-client-protocol = { workspace = true, features = ["unstable_mcp_over_acp"] } async-stream.workspace = true axum.workspace = true +base64.workspace = true futures.workspace = true -futures-concurrency.workspace = true -rustc-hash.workspace = true +hmac = "0.12" serde_json.workspace = true -thiserror = "2.0" +sha2 = "0.10" tokio = { workspace = true, features = ["net"] } tracing.workspace = true uuid.workspace = true [lints] workspace = true + +[dev-dependencies] +tokio = { workspace = true, features = ["io-util", "macros", "rt"] } diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs deleted file mode 100644 index 908b6b87..00000000 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/actor.rs +++ /dev/null @@ -1,78 +0,0 @@ -use agent_client_protocol::{ConnectTo, Dispatch, DynConnectTo, role::mcp}; -use futures::{SinkExt as _, StreamExt as _, channel::mpsc}; -use tracing::info; - -use super::BridgeMessage; - -/// Actor that bridges a single MCP connection between a local MCP client -/// and the ACP proxy chain. -#[derive(Debug)] -pub(crate) struct BridgeConnectionActor { - /// The loopback HTTP transport accepted by the compatibility listener. - transport: DynConnectTo, - - /// Sender for messages back to the polyfill's bridge runner loop. - bridge_tx: mpsc::Sender, - - /// Receiver for messages from the polyfill to forward to the MCP client. - to_mcp_client_rx: mpsc::Receiver, -} - -impl BridgeConnectionActor { - pub fn new( - component: impl ConnectTo, - bridge_tx: mpsc::Sender, - to_mcp_client_rx: mpsc::Receiver, - ) -> Self { - Self { - transport: DynConnectTo::new(component), - bridge_tx, - to_mcp_client_rx, - } - } - - pub async fn run(self, connection_id: String) -> Result<(), agent_client_protocol::Error> { - info!(connection_id, "MCP bridge connected"); - - let Self { - transport, - mut bridge_tx, - to_mcp_client_rx, - } = self; - - let result = mcp::Client - .builder() - .name(format!("mcp-client-to-polyfill({connection_id})")) - .on_receive_dispatch( - { - let mut bridge_tx = bridge_tx.clone(); - let connection_id = connection_id.clone(); - async move |message: Dispatch, _cx| { - bridge_tx - .send(BridgeMessage::ClientToServer { - connection_id: connection_id.clone(), - message, - }) - .await - .map_err(|_| agent_client_protocol::Error::internal_error()) - } - }, - agent_client_protocol::on_receive_dispatch!(), - ) - .connect_with(transport, async move |mcp_connection_to_client| { - let mut to_mcp_client_rx = to_mcp_client_rx; - while let Some(message) = to_mcp_client_rx.next().await { - mcp_connection_to_client.send_proxied_message(message)?; - } - Ok(()) - }) - .await; - - bridge_tx - .send(BridgeMessage::Disconnected { connection_id }) - .await - .map_err(|_| agent_client_protocol::Error::internal_error())?; - - result - } -} diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs index afed504b..47300394 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/http.rs @@ -1,1402 +1,948 @@ -//! HTTP-based MCP bridge transport. +//! MCP 2026-07-28 request-scoped Streamable HTTP endpoint. +//! +//! The adapter keeps raw MCP envelopes, ACP cancellation, and response-stream +//! lifetimes explicit. Each POST owns one operation, not an MCP session. -use agent_client_protocol::{ - BoxFuture, Channel, ConnectTo, RawJsonRpcMessage, RawJsonRpcParams, TransportBatchEntry, - TransportFrame, - role::mcp, - schema::v1::{ - Notification as RpcNotification, Request as RpcRequest, RequestId, Response as RpcResponse, - }, -}; +use std::{convert::Infallible, sync::Arc}; + +use agent_client_protocol::Error; use axum::{ - Router, - extract::State, - http::StatusCode, - response::{IntoResponse, Response, Sse}, - routing::post, + Json, Router, + body::{Body, HttpBody as _, to_bytes}, + extract::{Path, State}, + http::{HeaderMap, StatusCode, header}, + response::{ + IntoResponse, Response, Sse, + sse::{Event, KeepAlive}, + }, + routing::any, }; -use futures::{SinkExt, StreamExt as _, channel::mpsc, future::Either, stream::Stream}; -use futures_concurrency::future::FutureExt as _; -use futures_concurrency::stream::StreamExt as _; -use rustc_hash::FxHashMap; -use std::{ - collections::{HashMap, VecDeque}, - pin::pin, - sync::Arc, +use base64::Engine as _; +use futures::{SinkExt, StreamExt, channel::mpsc}; +use hmac::{Hmac, Mac}; +use serde_json::{Map, Value}; +use sha2::Sha256; +use tokio::{ + net::TcpListener, + sync::{Semaphore, mpsc as tokio_mpsc, oneshot}, }; -use tokio::net::TcpListener; -use super::{BridgeConnection, BridgeMessage, actor::BridgeConnectionActor}; +use super::BridgeMessage; -/// Runs an HTTP listener for MCP bridge connections. -pub async fn run_http_listener( - tcp_listener: TcpListener, - server_id: String, - mut bridge_tx: mpsc::Sender, -) -> Result<(), agent_client_protocol::Error> { - let (to_mcp_client_tx, to_mcp_client_rx) = mpsc::channel(128); +const VERSION: &str = "2026-07-28"; +const MAX_REQUEST_BODY_BYTES: usize = 1024 * 1024; - 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), - }) - .await - .map_err(|_| agent_client_protocol::Error::internal_error())?; +fn server_route(server_id: &str) -> String { + // Even an empty opaque ID must occupy a real route segment. + format!( + "mcp-{}", + base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(server_id) + ) +} - Ok(()) +pub(super) struct BridgeState { + secret: [u8; 32], + admission: Arc, + tx: mpsc::Sender, } -/// A component that receives HTTP requests/responses using the HTTP transport -/// defined by the MCP protocol. -struct HttpMcpBridge { - listener: tokio::net::TcpListener, +pub(super) async fn run_http_listener( + listener: TcpListener, + state: Arc, +) -> Result<(), Error> { + let app = Router::new() + .route("/{route}", any(handle_request)) + .with_state(state); + axum::serve(listener, app) + .await + .map_err(Error::into_internal_error) } -impl HttpMcpBridge { - /// Creates a new HTTP-MCP bridge from an existing TCP listener. - fn new(listener: tokio::net::TcpListener) -> Self { - Self { listener } +impl BridgeState { + pub(super) fn new(tx: mpsc::Sender) -> Arc { + Arc::new(Self { + secret: { + let mut secret = [0; 32]; + secret[..16].copy_from_slice(uuid::Uuid::new_v4().as_bytes()); + secret[16..].copy_from_slice(uuid::Uuid::new_v4().as_bytes()); + secret + }, + admission: Arc::new(Semaphore::new(super::MAX_ACTIVE_REQUESTS)), + tx, + }) } -} -impl ConnectTo for HttpMcpBridge { - async fn connect_to( - self, - client: impl ConnectTo, - ) -> Result<(), agent_client_protocol::Error> { - let (channel, serve_self) = self.into_channel_and_future(); - match futures::future::select(pin!(client.connect_to(channel)), serve_self).await { - Either::Left((result, _)) | Either::Right((result, _)) => result, - } + fn mac(&self, server_id: &str) -> Hmac { + let mut mac = Hmac::::new_from_slice(&self.secret).expect("SHA-256 HMAC key"); + mac.update(b"mcp-over-acp-http-adapter/server/v1\0"); + mac.update(server_id.as_bytes()); + mac } - fn into_channel_and_future( - self, - ) -> ( - Channel, - BoxFuture<'static, Result<(), agent_client_protocol::Error>>, - ) - where - Self: Sized, - { - let (channel_a, channel_b) = Channel::duplex(); - (channel_a, Box::pin(run(self.listener, channel_b))) + pub(super) fn declaration_url(&self, port: u16, server_id: &str) -> (String, String) { + let route = server_route(server_id); + let token = base64::engine::general_purpose::URL_SAFE_NO_PAD + .encode(self.mac(server_id).finalize().into_bytes()); + (format!("http://127.0.0.1:{port}/{route}"), token) } } -/// Error type for responding to malformed HTTP requests. -#[derive(Debug, thiserror::Error)] -#[error(transparent)] -struct HttpError(#[from] agent_client_protocol::Error); - -impl From for HttpError { - fn from(error: axum::Error) -> Self { - HttpError(agent_client_protocol::util::internal_error(error)) - } +fn error(status: StatusCode, id: Value, code: i64, message: &str) -> Response { + (status, Json(rpc_error(id, code, message))).into_response() } -impl IntoResponse for HttpError { - fn into_response(self) -> Response { - let message = format!("Error: {}", self.0); - (StatusCode::INTERNAL_SERVER_ERROR, message).into_response() +pub(super) fn rpc_error(id: Value, code: i64, message: &str) -> Value { + let mut response = + serde_json::json!({"jsonrpc":"2.0", "error":{"code":code,"message":message}}); + if !id.is_null() { + response["id"] = id; } + response } -/// 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)); +fn valid_request_id(id: &Value) -> bool { + id.is_string() || id.as_i64().is_some() || id.as_u64().is_some() +} - axum::serve(listener, app) - .await - .map_err(agent_client_protocol::util::internal_error) +/// Only the MCP 2026 payload metadata carries a subscription identifier. +/// Other fields (including opaque requestState and progress tokens) are untouched. +pub(super) fn rewrite_subscription_id(payload: &mut Value, request_id: &str, http_id: &Value) { + if let Some(subscription_id) = payload + .get_mut("_meta") + .and_then(Value::as_object_mut) + .and_then(|meta| meta.get_mut("io.modelcontextprotocol/subscriptionId")) + && subscription_id.as_str() == Some(request_id) + { + *subscription_id = http_id.clone(); } - .race(RunningServer::new().run(channel, registration_rx)) - .await } -/// The state we pass to our POST/GET handlers. -struct BridgeState { - /// Where to send registration messages. - registration_tx: mpsc::UnboundedSender, +pub(super) fn rpc_result(id: Value, request_id: &str, mut result: Value) -> Value { + rewrite_subscription_id(&mut result, request_id, &id); + serde_json::json!({"jsonrpc":"2.0", "id":id, "result":result}) } -/// Messages from HTTP handlers to the bridge server. -#[derive(Debug)] -#[allow(dead_code)] -enum HttpMessage { - /// A JSON-RPC request (has an id, expects a response via the channel). - Request { - http_request_id: uuid::Uuid, - request: RpcRequest, - response_tx: mpsc::UnboundedSender, - }, - /// A JSON-RPC notification (no id, no response expected). - Notification { - http_request_id: uuid::Uuid, - request: RpcNotification, - }, - /// A JSON-RPC response from the client. - Response { - http_request_id: uuid::Uuid, - response: RpcResponse, - }, - /// A batch retained as one transport frame. - Frame { - http_request_id: uuid::Uuid, - frame: TransportFrame, - request_ids: Vec, - response_tx: Option>, - }, - /// A GET request to open an SSE stream for server-initiated messages. - Get { - http_request_id: uuid::Uuid, - response_tx: mpsc::UnboundedSender, - }, +pub(super) fn rpc_binding_error(id: Value, error: Error) -> Value { + let value = serde_json::to_value(error).unwrap_or(Value::Null); + let peer_code = value.get("code").and_then(Value::as_i64); + let code = match peer_code { + Some(-33000 | -33001 | -33002 | -32800) => peer_code.unwrap(), + _ => -33002, + }; + let message = value + .get("message") + .and_then(Value::as_str) + .unwrap_or("MCP binding failure"); + rpc_error(id, code, message) } -struct RunningServer { - waiting_sessions: FxHashMap, - waiting_batch_sessions: Vec, - pending_calls: VecDeque, - general_sessions: Vec, - message_deque: VecDeque, +pub(super) fn rpc_peer_error(id: Value, error: Value) -> Value { + serde_json::json!({"jsonrpc":"2.0", "id":id, "error":error}) } -impl RunningServer { - fn new() -> Self { - RunningServer { - waiting_sessions: HashMap::default(), - waiting_batch_sessions: Vec::new(), - pending_calls: VecDeque::new(), - general_sessions: Vec::default(), - message_deque: VecDeque::with_capacity(32), - } - } - - /// The main loop: listen for incoming HTTP messages and outgoing JSON-RPC messages. - async fn run( - mut self, - mut channel: Channel, - http_rx: mpsc::UnboundedReceiver, - ) -> Result<(), agent_client_protocol::Error> { - #[derive(Debug)] - enum MultiplexMessage { - FromHttpToChannel(HttpMessage), - FromChannelToHttp(TransportFrame), - } - - let mut merged_stream = http_rx - .map(MultiplexMessage::FromHttpToChannel) - .merge(channel.rx.map(MultiplexMessage::FromChannelToHttp)); - - while let Some(message) = merged_stream.next().await { - tracing::trace!(?message, "received message"); - - match message { - MultiplexMessage::FromHttpToChannel(http_message) => { - self.handle_http_message(http_message, &mut channel.tx)?; - } - MultiplexMessage::FromChannelToHttp(message) => { - self.message_deque.push_back(message); - } - } +fn header_value<'a>(headers: &'a HeaderMap, name: &str) -> Option<&'a str> { + let mut values = headers.get_all(name).iter(); + let value = values.next()?.to_str().ok()?; + values.next().is_none().then_some(value) +} - self.drain_jsonrpc_messages(); - self.activate_pending_calls(&mut channel.tx)?; - } +fn valid_origin(headers: &HeaderMap) -> bool { + // Browsers supply Origin; only the actual loopback origin is trusted. + // Non-browser HTTP clients normally omit Origin. + let Some(origin) = header_value(headers, "origin") else { + return !headers.contains_key("origin"); + }; + let Some(host) = header_value(headers, "host") else { + return false; + }; + host.split_once(':') + .is_some_and(|(address, port)| address == "127.0.0.1" && port.parse::().is_ok()) + && origin == format!("http://{host}") +} - Ok(()) +/// HTTP qvalues are decimal 0..1 with at most three fractional digits, not +/// floating-point syntax (which also accepts NaN, exponents, and signs). +fn positive_quality(value: &str) -> Option { + let (whole, fraction) = value.split_once('.').unwrap_or((value, "")); + if fraction.len() > 3 || !fraction.bytes().all(|byte| byte.is_ascii_digit()) { + return None; } + match whole { + "0" => Some(fraction.bytes().any(|byte| byte != b'0')), + "1" if fraction.bytes().all(|byte| byte == b'0') => Some(true), + _ => None, + } +} - /// Handle an incoming HTTP message (request, notification, response, or GET). - fn handle_http_message( - &mut self, - message: HttpMessage, - channel_tx: &mut mpsc::UnboundedSender, - ) -> Result<(), agent_client_protocol::Error> { - match message { - HttpMessage::Request { - http_request_id, - request, - response_tx, - } => { - tracing::debug!(%http_request_id, ?request, "handling request"); - let request_id = request.id.clone(); - self.send_or_queue_call( - PendingCall { - frame: TransportFrame::Single(RawJsonRpcMessage::Request(request)), - request_ids: vec![request_id], - session: RegisteredSession::new(response_tx), - }, - channel_tx, - )?; - } - HttpMessage::Notification { - http_request_id: _, - request, - } => { - channel_tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::Notification( - request, - ))) - .map_err(agent_client_protocol::util::internal_error)?; - } - HttpMessage::Response { - http_request_id: _, - response, - } => { - channel_tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::Response( - response, - ))) - .map_err(agent_client_protocol::util::internal_error)?; - } - HttpMessage::Frame { - http_request_id, - frame, - request_ids, - response_tx, - } => { - tracing::debug!(%http_request_id, ?frame, "handling retained frame"); - if let Some(response_tx) = response_tx { - match &frame { - TransportFrame::Batch(_) => { - self.send_or_queue_call( - PendingCall { - frame, - request_ids, - session: RegisteredSession::new(response_tx), - }, - channel_tx, - )?; - return Ok(()); - } - TransportFrame::Single(_) | TransportFrame::Malformed { .. } => { - unreachable!("only batches use the retained frame variant") - } +fn accepts_both(headers: &HeaderMap) -> bool { + let mut json = false; + let mut sse = false; + for value in headers.get_all(header::ACCEPT) { + let Ok(value) = value.to_str() else { + return false; + }; + for item in value.split(',') { + let mut parts = item.split(';'); + let media = parts.next().unwrap_or("").trim(); + let mut quality = None; + for part in parts { + if let Some((key, q)) = part.trim().split_once('=') + && key.trim().eq_ignore_ascii_case("q") + { + let Some(positive) = positive_quality(q.trim()) else { + return false; + }; + if quality.replace(positive).is_some() { + return false; } } - channel_tx - .unbounded_send(frame) - .map_err(agent_client_protocol::util::internal_error)?; } - HttpMessage::Get { - http_request_id: _, - response_tx, - } => { - self.general_sessions - .push(RegisteredSession::new(response_tx)); + if quality == Some(false) { + continue; } + json |= media.eq_ignore_ascii_case("application/json"); + sse |= media.eq_ignore_ascii_case("text/event-stream"); } - self.purge_closed_sessions(); - Ok(()) - } - - fn send_or_queue_call( - &mut self, - call: PendingCall, - channel_tx: &mut mpsc::UnboundedSender, - ) -> Result<(), agent_client_protocol::Error> { - if self.call_conflicts_with_active(&call.request_ids) { - tracing::debug!( - request_ids = ?call.request_ids, - "queueing HTTP call until overlapping request IDs are no longer in flight" - ); - self.pending_calls.push_back(call); - return Ok(()); - } - - self.activate_call(call, channel_tx) } + json && sse +} - fn activate_call( - &mut self, - call: PendingCall, - channel_tx: &mut mpsc::UnboundedSender, - ) -> Result<(), agent_client_protocol::Error> { - let PendingCall { - frame, - request_ids, - session, - } = call; - let is_batch = matches!(frame, TransportFrame::Batch(_)); - channel_tx - .unbounded_send(frame) - .map_err(agent_client_protocol::util::internal_error)?; - - if is_batch { - self.waiting_batch_sessions.push(WaitingBatchSession { - request_ids, - session, - }); - } else { - let request_id = request_ids - .into_iter() - .next() - .expect("single request calls always have one request ID"); - self.waiting_sessions.insert(request_id, session); - } - - Ok(()) +fn mirrored_name<'a>(method: &str, params: &'a Map) -> Option<&'a str> { + match method { + "tools/call" | "prompts/get" => params.get("name").and_then(Value::as_str), + "resources/read" => params.get("uri").and_then(Value::as_str), + _ => None, } +} - fn activate_pending_calls( - &mut self, - channel_tx: &mut mpsc::UnboundedSender, - ) -> Result<(), agent_client_protocol::Error> { - loop { - let Some(call) = self.pending_calls.front() else { - return Ok(()); - }; - if call.session.outgoing_tx.is_closed() { - self.pending_calls.pop_front(); - continue; - } - if self.call_conflicts_with_active(&call.request_ids) { - return Ok(()); - } - - let call = self - .pending_calls - .pop_front() - .expect("pending call was checked above"); - self.activate_call(call, channel_tx)?; - } +/// Decode the MCP sentinel; rejecting invalid or noncanonical Base64 prevents +/// intermediaries and the adapter from disagreeing on mirrored routing values. +fn matches_mirror(header: Option<&str>, body: &str) -> bool { + let Some(header) = header else { + return false; + }; + if let Some(encoded) = header + .strip_prefix("=?base64?") + .and_then(|h| h.strip_suffix("?=")) + { + base64::engine::general_purpose::STANDARD + .decode(encoded) + .is_ok_and(|bytes| bytes == body.as_bytes()) + } else { + // Literal sentinel-looking values must be encoded to avoid ambiguity. + !(header.starts_with("=?base64?") && header.ends_with("?=")) && header == body } +} - fn call_conflicts_with_active(&self, request_ids: &[RequestId]) -> bool { - let unidentified_batch_is_active = self - .waiting_batch_sessions - .iter() - .any(|waiting| waiting.request_ids.is_empty()); - - if request_ids.is_empty() { - // A response-bearing batch without a request ID (for example, a - // notification plus an invalid scalar) receives a grouped error - // response whose only ID is null. Keep those responses ordered, - // including with explicit null-ID calls, because the wire response - // does not otherwise carry enough provenance to distinguish them. - return unidentified_batch_is_active || self.request_id_is_active(&RequestId::Null); - } - - request_ids - .iter() - .any(|request_id| self.request_id_is_active(request_id)) - || unidentified_batch_is_active && request_ids.contains(&RequestId::Null) +async fn handle_request( + State(state): State>, + Path(route): Path, + method: axum::http::Method, + headers: HeaderMap, + body: Body, +) -> Response { + if [ + "origin", + "authorization", + "mcp-protocol-version", + "mcp-method", + "mcp-name", + ] + .into_iter() + .any(|name| headers.get_all(name).iter().nth(1).is_some()) + { + return error( + StatusCode::BAD_REQUEST, + Value::Null, + -32020, + "HeaderMismatch: duplicate routing or authentication header", + ); } - - fn request_id_is_active(&self, request_id: &RequestId) -> bool { - self.waiting_sessions.contains_key(request_id) - || self - .waiting_batch_sessions - .iter() - .any(|waiting| waiting.request_ids.contains(request_id)) + if !valid_origin(&headers) { + return error(StatusCode::FORBIDDEN, Value::Null, -32600, "Invalid Origin"); } - - fn drain_jsonrpc_messages(&mut self) { - while let Some(message) = self.message_deque.pop_front() { - if let Some(message) = self.try_dispatch_jsonrpc_message(message) { - self.message_deque.push_front(message); - break; - } - } + if route.len() > 4096 { + return error( + StatusCode::NOT_FOUND, + Value::Null, + -32601, + "Unknown MCP route", + ); } - - fn try_dispatch_jsonrpc_message( - &mut self, - mut message: TransportFrame, - ) -> Option { - if matches!(message, TransportFrame::Malformed { .. }) { - // Malformed frames emitted by a relay are wire data, not protocol - // responses, so they are delivered through a general stream. - } else if matches!(message, TransportFrame::Batch(_)) { - let response_ids: Vec<_> = match &message { - TransportFrame::Batch(batch) => batch - .entries() - .filter_map(|entry| match entry { - TransportBatchEntry::Message(message) => message.response_id(), - TransportBatchEntry::Malformed { .. } => None, - }) - .collect(), - _ => unreachable!(), - }; - let correlated = self.waiting_batch_sessions.iter().position(|waiting| { - !waiting.request_ids.is_empty() - && waiting - .request_ids - .iter() - .any(|id| response_ids.contains(&id)) - }); - let fallback = response_ids.contains(&&RequestId::Null).then(|| { - self.waiting_batch_sessions - .iter() - .position(|waiting| waiting.request_ids.is_empty()) - }); - let fallback = fallback.flatten(); - if let Some(index) = correlated.or(fallback) { - let session = self.waiting_batch_sessions.remove(index).session; - // This response belongs to that HTTP POST even if its SSE - // receiver has gone away. Never let it fall through to a - // later request that reuses the same JSON-RPC ID. - drop(session.outgoing_tx.unbounded_send(message)); - return None; - } - } - - let message_id = match &message { - TransportFrame::Single(message) => message.response_id().cloned(), - TransportFrame::Malformed { .. } => None, - TransportFrame::Batch(batch) => batch.entries().find_map(|entry| match entry { - TransportBatchEntry::Message(message) => message.response_id().cloned(), - TransportBatchEntry::Malformed { .. } => None, - }), + let server_id = { + let decoded = route.strip_prefix("mcp-").and_then(|encoded| { + base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(encoded) + .ok() + }); + let Some(server_id) = decoded + .and_then(|id| String::from_utf8(id).ok()) + .filter(|id| server_route(id) == route) + else { + return error( + StatusCode::NOT_FOUND, + Value::Null, + -32601, + "Unknown MCP route", + ); }; - - if let Some(ref message_id) = message_id - && let Some(session) = self.waiting_sessions.remove(message_id) - { - // This response belongs to that HTTP POST even if its SSE - // receiver has gone away. Never let it fall through to a later - // request that reuses the same JSON-RPC ID. - drop(session.outgoing_tx.unbounded_send(message)); - return None; - } - - self.purge_closed_sessions(); - let all_sessions = self - .general_sessions - .iter_mut() - .chain(self.waiting_sessions.values_mut()) - .chain( - self.waiting_batch_sessions - .iter_mut() - .map(|waiting| &mut waiting.session), - ) - .chain( - self.pending_calls - .iter_mut() - .map(|waiting| &mut waiting.session), + let authorization = + header_value(&headers, "authorization").and_then(|value| value.split_once(' ')); + if !authorization.is_some_and(|(scheme, supplied)| { + scheme.eq_ignore_ascii_case("bearer") + && base64::engine::general_purpose::URL_SAFE_NO_PAD + .decode(supplied) + .is_ok_and(|tag| state.mac(&server_id).verify_slice(&tag).is_ok()) + }) { + let mut response = error( + StatusCode::UNAUTHORIZED, + Value::Null, + -32600, + "Unauthorized", ); - for session in all_sessions { - match session.outgoing_tx.unbounded_send(message) { - Ok(()) => return None, - Err(m) => { - assert!(m.is_disconnected()); - message = m.into_inner(); - } - } + response.headers_mut().insert( + header::WWW_AUTHENTICATE, + "Bearer".parse().expect("static header"), + ); + return response; } - - Some(message) + server_id + }; + if method != axum::http::Method::POST { + return error( + StatusCode::METHOD_NOT_ALLOWED, + Value::Null, + -32600, + "Only POST is supported", + ); } - - fn purge_closed_sessions(&mut self) { - self.general_sessions - .retain(|session| !session.outgoing_tx.is_closed()); - self.pending_calls - .retain(|call| !call.session.outgoing_tx.is_closed()); - - // Calls already forwarded to the JSON-RPC peer stay registered until - // their response arrives. Otherwise a late response could be routed - // to a newer HTTP POST that reused the same request ID. + if !accepts_both(&headers) { + return error( + StatusCode::NOT_ACCEPTABLE, + Value::Null, + -32600, + "Accept must include application/json and text/event-stream", + ); } -} - -struct PendingCall { - frame: TransportFrame, - request_ids: Vec, - session: RegisteredSession, -} - -struct WaitingBatchSession { - request_ids: Vec, - session: RegisteredSession, -} - -struct RegisteredSession { - #[allow(dead_code)] - id: uuid::Uuid, - outgoing_tx: mpsc::UnboundedSender, -} - -impl RegisteredSession { - fn new(outgoing_tx: mpsc::UnboundedSender) -> Self { - Self { - id: uuid::Uuid::new_v4(), - outgoing_tx, - } + if header_value(&headers, header::CONTENT_TYPE.as_str()).is_none_or(|value| { + !value + .split(';') + .next() + .is_some_and(|media| media.trim().eq_ignore_ascii_case("application/json")) + }) { + return error( + StatusCode::UNSUPPORTED_MEDIA_TYPE, + Value::Null, + -32600, + "Expected application/json", + ); } -} - -/// Accept a POST request carrying a JSON-RPC frame from an MCP client. -/// For response-bearing calls and batches, we return an SSE stream. For -/// notification/response-only frames, we return 202 Accepted. -async fn handle_post( - State(state): State>, - body: String, -) -> Result { - let http_request_id = uuid::Uuid::new_v4(); - let frame = TransportFrame::parse_json(&body); - - match frame { - TransportFrame::Single(message) => match message { - RawJsonRpcMessage::Request(request) => { - let (tx, rx) = mpsc::unbounded(); - state - .registration_tx - .unbounded_send(HttpMessage::Request { - http_request_id, - request, - response_tx: tx, - }) - .map_err(agent_client_protocol::util::internal_error)?; - - Ok(sse_response(rx)) - } - RawJsonRpcMessage::Notification(request) => { - state - .registration_tx - .unbounded_send(HttpMessage::Notification { - http_request_id, - request, - }) - .map_err(agent_client_protocol::util::internal_error)?; - Ok(StatusCode::ACCEPTED.into_response()) - } - RawJsonRpcMessage::Response(response) => { - state - .registration_tx - .unbounded_send(HttpMessage::Response { - http_request_id, - response, - }) - .map_err(agent_client_protocol::util::internal_error)?; - Ok(StatusCode::ACCEPTED.into_response()) - } - }, - TransportFrame::Malformed { raw, error } => { - if raw - .parse::() - .is_ok_and(|value| is_response_only_shape(&value)) - { - return Ok(StatusCode::ACCEPTED.into_response()); - } - Ok(immediate_sse_response(TransportFrame::Single( - RawJsonRpcMessage::response(RequestId::Null, Err(error)), - ))) - } - TransportFrame::Batch(batch) => { - if batch - .entries() - .all(|entry| matches!(entry, TransportBatchEntry::Malformed { .. })) - { - let responses = agent_client_protocol::TransportBatch::from_messages( - batch.entries().filter_map(|entry| { - let TransportBatchEntry::Malformed { raw, error } = entry else { - unreachable!("all batch entries were checked as malformed") - }; - (!is_response_only_shape(raw)).then(|| { - RawJsonRpcMessage::response(RequestId::Null, Err(error.clone())) - }) - }), - ); - let Some(responses) = responses else { - return Ok(StatusCode::ACCEPTED.into_response()); - }; - return Ok(immediate_sse_response(TransportFrame::Batch(responses))); - } - - let mut request_ids = Vec::new(); - let mut expects_response = false; - for entry in batch.entries() { - match entry { - TransportBatchEntry::Message(RawJsonRpcMessage::Request(request)) => { - request_ids.push(request.id.clone()); - expects_response = true; - } - TransportBatchEntry::Malformed { raw, .. } => { - expects_response |= !is_response_only_shape(raw); - } - TransportBatchEntry::Message( - RawJsonRpcMessage::Notification(_) | RawJsonRpcMessage::Response(_), - ) => {} - } - } - let frame = TransportFrame::Batch(batch); - if expects_response { - let (tx, rx) = mpsc::unbounded(); - state - .registration_tx - .unbounded_send(HttpMessage::Frame { - http_request_id, - frame, - request_ids, - response_tx: Some(tx), - }) - .map_err(agent_client_protocol::util::internal_error)?; - Ok(sse_response(rx)) - } else { - state - .registration_tx - .unbounded_send(HttpMessage::Frame { - http_request_id, - frame, - request_ids, - response_tx: None, - }) - .map_err(agent_client_protocol::util::internal_error)?; - Ok(StatusCode::ACCEPTED.into_response()) - } - } + // Acquire before reading a potentially slow/large request body. The permit + // stays owned by the response body until the client consumes or drops it. + let Ok(permit) = state.admission.clone().try_acquire_owned() else { + return error( + StatusCode::TOO_MANY_REQUESTS, + Value::Null, + -33000, + "Too many outstanding MCP responses", + ); + }; + let response = handle_admitted_request(state, server_id, headers, body).await; + let (mut parts, body) = response.into_parts(); + if let Some(length) = body.size_hint().exact() { + parts + .headers + .entry(header::CONTENT_LENGTH) + .or_insert_with(|| length.to_string().parse().expect("decimal body length")); } -} - -fn is_response_only_shape(value: &serde_json::Value) -> bool { - value.as_object().is_some_and(|object| { - !object.contains_key("method") - && (object.contains_key("result") || object.contains_key("error")) - }) -} - -/// Accept a GET request from an MCP client. -/// Opens an SSE stream for server-initiated messages. -async fn handle_get( - State(state): State>, -) -> Result>>, HttpError> { - let http_request_id = uuid::Uuid::new_v4(); - let (tx, mut rx) = mpsc::unbounded(); - state - .registration_tx - .unbounded_send(HttpMessage::Get { - http_request_id, - response_tx: tx, - }) - .map_err(agent_client_protocol::util::internal_error)?; - + // One ownership rule for every admitted response, including validation + // failures that echo a potentially large, but valid, external request ID. let stream = async_stream::stream! { - while let Some(message) = rx.next().await { - yield sse_event(message); + let _permit = permit; + let mut body = body.into_data_stream(); + while let Some(chunk) = body.next().await { + yield chunk; } }; - - Ok(Sse::new(stream)) -} - -fn sse_event(frame: TransportFrame) -> Result { - Ok(axum::response::sse::Event::default().data(frame.to_json()?)) + Response::from_parts(parts, Body::from_stream(stream)) } -fn sse_response(mut rx: mpsc::UnboundedReceiver) -> Response { +async fn handle_admitted_request( + state: Arc, + server_id: String, + headers: HeaderMap, + body: Body, +) -> Response { + let Ok(body) = to_bytes(body, MAX_REQUEST_BODY_BYTES).await else { + return error( + StatusCode::PAYLOAD_TOO_LARGE, + Value::Null, + -33000, + "Request body too large", + ); + }; + let body: Value = match serde_json::from_slice(&body) { + Ok(body) => body, + Err(_) => return error(StatusCode::BAD_REQUEST, Value::Null, -32700, "Parse error"), + }; + let id = body + .get("id") + .filter(|id| valid_request_id(id)) + .cloned() + .unwrap_or(Value::Null); + let Some(object) = body.as_object() else { + return error( + StatusCode::BAD_REQUEST, + id, + -32600, + "Expected one JSON-RPC request", + ); + }; + if object.get("jsonrpc").and_then(Value::as_str) != Some("2.0") + || object.contains_key("result") + || object.contains_key("error") + || object.get("id").is_none_or(|id| !valid_request_id(id)) + { + return error( + StatusCode::BAD_REQUEST, + id, + -32600, + "Expected one JSON-RPC request; batches, notifications and client responses are unsupported", + ); + } + let Some(method) = object.get("method").and_then(Value::as_str) else { + return error(StatusCode::BAD_REQUEST, id, -32600, "Missing method"); + }; + let params = match object.get("params") { + None => None, + Some(Value::Object(params)) => Some(params.clone()), + _ => { + return error( + StatusCode::BAD_REQUEST, + id, + -32602, + "Parameters must be an object", + ); + } + }; + let metadata_version = params + .as_ref() + .and_then(|params| params.get("_meta")) + .and_then(Value::as_object) + .and_then(|meta| meta.get("io.modelcontextprotocol/protocolVersion")) + .and_then(Value::as_str); + let version = header_value(&headers, "mcp-protocol-version"); + if version.is_none() || version != metadata_version { + return error( + StatusCode::BAD_REQUEST, + id, + -32020, + "HeaderMismatch: MCP-Protocol-Version does not match params._meta", + ); + } + if version != Some(VERSION) { + return ( + StatusCode::BAD_REQUEST, + Json(serde_json::json!({ + "jsonrpc":"2.0","id":id, + "error":{"code":-32022,"message":"Unsupported protocol version", + "data":{"supported":[VERSION],"requested":version}} + })), + ) + .into_response(); + } + if header_value(&headers, "mcp-method") != Some(method) { + return error( + StatusCode::BAD_REQUEST, + id, + -32020, + "HeaderMismatch: Mcp-Method does not match method", + ); + } + if matches!(method, "tools/call" | "prompts/get" | "resources/read") { + let Some(name) = params + .as_ref() + .and_then(|params| mirrored_name(method, params)) + else { + return error( + StatusCode::BAD_REQUEST, + id, + -32602, + "Missing params.name or params.uri", + ); + }; + if !matches_mirror(header_value(&headers, "mcp-name"), name) { + return error( + StatusCode::BAD_REQUEST, + id, + -32020, + "HeaderMismatch: Mcp-Name does not match request", + ); + } + } + // This endpoint re-exports native tools without transport-only x-mcp-header + // annotations. Mirrored parameter headers have no authority here. + if headers + .keys() + .any(|key| key.as_str().starts_with("mcp-param-")) + { + return error( + StatusCode::BAD_REQUEST, + id, + -32020, + "HeaderMismatch: Mcp-Param headers are not supported by this adapter", + ); + } + if method == "initialize" || method.starts_with("notifications/") { + return error(StatusCode::NOT_FOUND, id, -32601, "Method not found"); + } + let id_for_bridge_error = id.clone(); + let (notification_tx, mut response_rx) = tokio_mpsc::channel(super::MAX_QUEUED_NOTIFICATIONS); + let response_tx = super::StreamSender { + tx: notification_tx, + used: Arc::new(std::sync::atomic::AtomicUsize::new(0)), + }; + let (terminal_tx, mut terminal_rx) = oneshot::channel(); + let message = BridgeMessage::Request { + server_id, + request_id: uuid::Uuid::new_v4().to_string(), + http_id: id, + method: method.into(), + params, + response_tx, + terminal_tx, + }; + let mut tx = state.tx.clone(); + if tx.send(message).await.is_err() { + return error( + StatusCode::SERVICE_UNAVAILABLE, + id_for_bridge_error.clone(), + -33002, + "ACP bridge unavailable", + ); + } + let first = tokio::select! { + biased; + notification = response_rx.recv(), if !response_rx.is_closed() || !response_rx.is_empty() => + match notification { + Some(mut message) => Some(std::mem::take(&mut message.value)), + None => (&mut terminal_rx).await.ok(), + }, + terminal = &mut terminal_rx => terminal.ok(), + }; + let Some(first) = first else { + return error( + StatusCode::SERVICE_UNAVAILABLE, + id_for_bridge_error, + -33002, + "ACP bridge closed", + ); + }; + if first.get("id").is_some() { + let status = if first.pointer("/error/code").and_then(Value::as_i64) == Some(-32601) { + StatusCode::NOT_FOUND + } else { + StatusCode::OK + }; + let payload = first.to_string(); + let length = payload.len().to_string(); + let stream = async_stream::stream! { + yield Ok::<_, Infallible>(axum::body::Bytes::from(payload)); + }; + return ( + status, + [ + (header::CONTENT_TYPE, "application/json".to_string()), + (header::CONTENT_LENGTH, length), + ], + Body::from_stream(stream), + ) + .into_response(); + } let stream = async_stream::stream! { - while let Some(message) = rx.next().await { - yield sse_event(message); + yield Ok::<_, Infallible>(Event::default().data(first.to_string())); + loop { + // Drain already-queued notifications before a successful final response. + // Overflow is delivered through the independent terminal path. + let message = tokio::select! { + biased; + notification = response_rx.recv(), if !response_rx.is_closed() || !response_rx.is_empty() => + match notification { + Some(mut message) => Some(std::mem::take(&mut message.value)), + None => (&mut terminal_rx).await.ok(), + }, + terminal = &mut terminal_rx => terminal.ok(), + }; + let Some(message) = message else { break }; + let final_response = message.get("id").is_some(); + yield Ok::<_, Infallible>(Event::default().data(message.to_string())); + if final_response { break } } }; - Sse::new(stream).into_response() -} - -fn immediate_sse_response(frame: TransportFrame) -> Response { - Sse::new(futures::stream::once(async move { sse_event(frame) })).into_response() + let mut response = Sse::new(stream) + .keep_alive(KeepAlive::default()) + .into_response(); + response + .headers_mut() + .insert("x-accel-buffering", "no".parse().expect("static header")); + response } #[cfg(test)] mod tests { use super::*; - - async fn single_sse_payload(response: Response) -> serde_json::Value { - let body = axum::body::to_bytes(response.into_body(), 64 * 1024) - .await - .expect("SSE response body"); - let body = std::str::from_utf8(&body).expect("UTF-8 SSE response"); - let payload = body - .lines() - .find_map(|line| line.strip_prefix("data:").map(str::trim_start)) - .expect("one SSE data event"); - serde_json::from_str(payload).expect("JSON-RPC SSE payload") - } - - async fn single_sse_message(response: Response) -> RawJsonRpcMessage { - serde_json::from_value(single_sse_payload(response).await) - .expect("single JSON-RPC SSE message") - } + use tokio::io::{AsyncReadExt, AsyncWriteExt}; #[test] - fn malformed_post_cannot_steal_a_valid_null_id_response() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); + fn stateless_declarations_do_not_allocate_routes() { + let (tx, _rx) = mpsc::channel(1); + let state = BridgeState::new(tx); + let first = state.declaration_url(1234, "server/one"); + let other = state.declaration_url(1234, "server/two"); + assert_eq!(first, state.declaration_url(1234, "server/one")); + assert_ne!(first, other); + assert!(state.declaration_url(1234, "").0.ends_with("/mcp-")); + for i in 0..1000 { + let (url, bearer) = state.declaration_url(1234, &i.to_string()); + assert!(url.starts_with("http://127.0.0.1:1234/")); + assert!(!url.contains(&bearer)); + } + } - let valid_http_response = handle_post( - State(state.clone()), - serde_json::json!({ - "jsonrpc": "2.0", - "id": null, - "method": "example", - "params": {} - }) + #[tokio::test] + async fn bridge_failure_preserves_valid_external_id() { + let (tx, rx) = mpsc::channel(1); + let state = BridgeState::new(tx); + drop(rx); + let (_, token) = state.declaration_url(8000, "server"); + let mut headers = HeaderMap::new(); + headers.insert("host", "127.0.0.1:8000".parse().unwrap()); + headers.insert("authorization", format!("Bearer {token}").parse().unwrap()); + headers.insert( + "accept", + "application/json, text/event-stream".parse().unwrap(), + ); + headers.insert("content-type", "application/json".parse().unwrap()); + headers.insert("mcp-protocol-version", VERSION.parse().unwrap()); + headers.insert("mcp-method", "tools/list".parse().unwrap()); + let response = handle_request( + State(state), + Path(server_route("server")), + axum::http::Method::POST, + headers, + Body::from( + serde_json::json!({"jsonrpc":"2.0","id":"external", + "method":"tools/list","params":{"_meta":{ + "io.modelcontextprotocol/protocolVersion":VERSION, + "io.modelcontextprotocol/clientCapabilities":{}}}}) .to_string(), - ) + ), + ) + .await; + assert_eq!(response.status(), StatusCode::SERVICE_UNAVAILABLE); + let bytes = to_bytes(response.into_body(), MAX_REQUEST_BODY_BYTES) .await - .expect("valid null-ID POST"); - let malformed_http_response = handle_post(State(state), "{not json".to_owned()) - .await - .expect("malformed POST receives a JSON-RPC error"); - - let valid_request = registration_rx - .next() - .await - .expect("valid request is forwarded"); - assert!( - registration_rx.try_recv().is_err(), - "malformed input must be answered by its own HTTP request" - ); - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - server - .handle_http_message(valid_request, &mut channel_tx) - .expect("forward valid request"); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(_))) - )); - - let valid_response = TransportFrame::Single(RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "valid" })), - )); - assert!( - server - .try_dispatch_jsonrpc_message(valid_response) - .is_none() - ); - - assert!(matches!( - single_sse_message(valid_http_response).await, - RawJsonRpcMessage::Response(RpcResponse::Result { - id: RequestId::Null, - result, - .. - }) if result == serde_json::json!({ "source": "valid" }) - )); - assert!(matches!( - single_sse_message(malformed_http_response).await, - RawJsonRpcMessage::Response(RpcResponse::Error { - id: RequestId::Null, - error, - .. - }) if error.code == agent_client_protocol::ErrorCode::ParseError - )); - }); + .unwrap(); + let body: Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(body["id"], "external"); + assert_eq!(body["error"]["code"], -33002); } - #[test] - fn malformed_response_shaped_posts_are_ignored_without_registration() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - let malformed_response = serde_json::json!({ - "jsonrpc": "2.0", - "id": 1, - "result": null, - "error": { "code": -32603, "message": "Internal error" } - }); - - let single = handle_post(State(state.clone()), malformed_response.to_string()) - .await - .expect("malformed response-shaped POST"); - assert_eq!(single.status(), StatusCode::ACCEPTED); - - let batch = handle_post( - State(state), - serde_json::Value::Array(vec![malformed_response]).to_string(), + #[tokio::test] + async fn unread_validation_errors_hold_admission_until_consumed_or_dropped() { + let (tx, _rx) = mpsc::channel(1); + let state = BridgeState::new(tx); + let (_, token) = state.declaration_url(8000, "server"); + let route = server_route("server"); + let mut headers = HeaderMap::new(); + headers.insert("authorization", format!("Bearer {token}").parse().unwrap()); + headers.insert( + "accept", + "application/json, text/event-stream".parse().unwrap(), + ); + headers.insert("content-type", "application/json".parse().unwrap()); + // No method: validation must echo this large known ID without releasing + // the permit while the client still owns its unread response. + let id = "external".repeat(32 * 1024); + let body = serde_json::json!({"jsonrpc":"2.0", "id":id}).to_string(); + let send = || { + handle_request( + State(state.clone()), + Path(route.clone()), + axum::http::Method::POST, + headers.clone(), + Body::from(body.clone()), ) - .await - .expect("malformed response-only batch POST"); - assert_eq!(batch.status(), StatusCode::ACCEPTED); - assert!( - registration_rx.try_recv().is_err(), - "ignored responses must not be forwarded or register HTTP waiters" - ); - }); - } + }; + let mut responses = Vec::new(); + for _ in 0..super::super::MAX_ACTIVE_REQUESTS { + let response = send().await; + assert_eq!(response.status(), StatusCode::BAD_REQUEST); + responses.push(response); + } + assert_eq!(send().await.status(), StatusCode::TOO_MANY_REQUESTS); - #[test] - fn malformed_response_sibling_does_not_hide_invalid_batch_value() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - let response = handle_post( - State(state), - serde_json::json!([ - 17, - { - "jsonrpc": "2.0", - "id": 1, - "result": null, - "error": { "code": -32603, "message": "Internal error" } - } - ]) - .to_string(), - ) + let bytes = to_bytes(responses.pop().unwrap().into_body(), MAX_REQUEST_BODY_BYTES) .await - .expect("mixed malformed batch POST"); - - let payload = single_sse_payload(response).await; - let entries = payload.as_array().expect("batch response array"); - assert_eq!(entries.len(), 1); - assert_eq!(entries[0]["id"], serde_json::Value::Null); - assert_eq!( - entries[0]["error"]["code"], - i32::from(agent_client_protocol::ErrorCode::InvalidRequest) - ); - assert!( - registration_rx.try_recv().is_err(), - "an all-malformed batch is answered by its originating POST" - ); - }); + .unwrap(); + let error: Value = serde_json::from_slice(&bytes).unwrap(); + assert_eq!(error["id"], id); + assert_eq!(error["error"]["code"], -32600); + assert_eq!(state.admission.available_permits(), 1); + responses.push(send().await); + assert_eq!(state.admission.available_permits(), 0); + + drop(responses.pop()); + assert_eq!(state.admission.available_permits(), 1); + let recovered = send().await; + assert_eq!(recovered.status(), StatusCode::BAD_REQUEST); + drop(recovered); + drop(responses); + assert_eq!( + state.admission.available_permits(), + super::super::MAX_ACTIVE_REQUESTS + ); } - #[test] - fn concurrent_null_id_posts_are_serialized_to_preserve_response_provenance() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - - let first_http_response = handle_post( + #[tokio::test] + async fn unread_terminal_bodies_hold_admission_until_drop() { + let (tx, mut rx) = mpsc::channel(128); + let state = BridgeState::new(tx); + let (_, token) = state.declaration_url(8000, "server"); + let route = server_route("server"); + let mut headers = HeaderMap::new(); + headers.insert("host", "127.0.0.1:8000".parse().unwrap()); + headers.insert("authorization", format!("bearer {token}").parse().unwrap()); + headers.insert("accept", "application/json".parse().unwrap()); + headers.append("accept", "text/event-stream;q=0.8".parse().unwrap()); + headers.insert( + "content-type", + "application/json; charset=utf-8".parse().unwrap(), + ); + headers.insert("mcp-protocol-version", VERSION.parse().unwrap()); + headers.insert("mcp-method", "tools/list".parse().unwrap()); + tokio::spawn(async move { + while let Some(BridgeMessage::Request { + terminal_tx, + http_id, + .. + }) = rx.next().await + { + drop(terminal_tx.send(rpc_result(http_id, "", serde_json::json!({"tools":[]})))); + } + }); + let body = serde_json::json!({"jsonrpc":"2.0","id":"known","method":"tools/list", + "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":VERSION, + "io.modelcontextprotocol/clientCapabilities":{}}}}) + .to_string(); + let send = || { + handle_request( State(state.clone()), - serde_json::json!({ - "jsonrpc": "2.0", - "id": null, - "method": "example/first", - "params": {} - }) - .to_string(), + Path(route.clone()), + axum::http::Method::POST, + headers.clone(), + Body::from(body.clone()), ) - .await - .expect("first null-ID POST"); - let second_http_response = handle_post( - State(state), - serde_json::json!({ - "jsonrpc": "2.0", - "id": null, - "method": "example/second", - "params": {} - }) - .to_string(), - ) - .await - .expect("second null-ID POST"); - - let first_registration = registration_rx.next().await.unwrap(); - let second_registration = registration_rx.next().await.unwrap(); - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - - server - .handle_http_message(first_registration, &mut channel_tx) - .unwrap(); - server - .handle_http_message(second_registration, &mut channel_tx) - .unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) - if request.method.as_ref() == "example/first" - && request.id == RequestId::Null - )); - assert!( - channel_rx.try_recv().is_err(), - "an overlapping null-ID request must wait for the first response" - ); - - assert!( - server - .try_dispatch_jsonrpc_message(TransportFrame::Single( - RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "first" })), - ), - )) - .is_none() - ); - server.activate_pending_calls(&mut channel_tx).unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) - if request.method.as_ref() == "example/second" - && request.id == RequestId::Null - )); - - assert!( - server - .try_dispatch_jsonrpc_message(TransportFrame::Single( - RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "second" })), - ), - )) - .is_none() - ); - - assert!(matches!( - single_sse_message(first_http_response).await, - RawJsonRpcMessage::Response(RpcResponse::Result { - id: RequestId::Null, - result, - .. - }) if result == serde_json::json!({ "source": "first" }) - )); - assert!(matches!( - single_sse_message(second_http_response).await, - RawJsonRpcMessage::Response(RpcResponse::Result { - id: RequestId::Null, - result, - .. - }) if result == serde_json::json!({ "source": "second" }) - )); - }); + }; + let mut responses = Vec::new(); + for _ in 0..super::super::MAX_ACTIVE_REQUESTS { + let response = send().await; + assert_eq!(response.status(), StatusCode::OK); + responses.push(response); + } + assert_eq!(send().await.status(), StatusCode::TOO_MANY_REQUESTS); + drop(responses.pop()); + let recovered = send().await; + assert_eq!(recovered.status(), StatusCode::OK); } #[test] - fn unidentified_batch_posts_are_serialized_to_preserve_response_provenance() { - futures::executor::block_on(async { - fn unidentified_batch(method: &str) -> TransportFrame { - TransportFrame::parse_json( - &serde_json::json!([ - { - "jsonrpc": "2.0", - "method": method, - "params": {} - }, - 17 - ]) - .to_string(), - ) - } - - fn grouped_response(source: &str) -> TransportFrame { - TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": source })), - ), - ]) - .expect("grouped response is non-empty"), - ) - } - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - let (first_tx, mut first_rx) = mpsc::unbounded(); - let (second_tx, mut second_rx) = mpsc::unbounded(); - - let first_frame = unidentified_batch("example/first"); - let second_frame = unidentified_batch("example/second"); - let expected_first_frame = first_frame.to_json().unwrap(); - let expected_second_frame = second_frame.to_json().unwrap(); - server - .handle_http_message( - HttpMessage::Frame { - http_request_id: uuid::Uuid::new_v4(), - frame: first_frame, - request_ids: Vec::new(), - response_tx: Some(first_tx), - }, - &mut channel_tx, - ) - .unwrap(); - server - .handle_http_message( - HttpMessage::Frame { - http_request_id: uuid::Uuid::new_v4(), - frame: second_frame, - request_ids: Vec::new(), - response_tx: Some(second_tx), - }, - &mut channel_tx, - ) - .unwrap(); - - assert_eq!( - channel_rx.next().await.unwrap().to_json().unwrap(), - expected_first_frame - ); - assert!( - channel_rx.try_recv().is_err(), - "a second unidentified batch must wait for the first response" - ); - - let callback = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::notification("callback".into(), serde_json::json!({})) - .unwrap(), - ]) - .expect("callback batch is non-empty"), - ); - assert!(server.try_dispatch_jsonrpc_message(callback).is_none()); - assert!(matches!( - first_rx.next().await, - Some(TransportFrame::Batch(_)) - )); - - let first_response = grouped_response("first"); - let expected_first_response = first_response.to_json().unwrap(); - assert!( - server - .try_dispatch_jsonrpc_message(first_response) - .is_none() - ); - assert_eq!( - first_rx.next().await.unwrap().to_json().unwrap(), - expected_first_response - ); - - server.activate_pending_calls(&mut channel_tx).unwrap(); - assert_eq!( - channel_rx.next().await.unwrap().to_json().unwrap(), - expected_second_frame - ); - - let second_response = grouped_response("second"); - let expected_second_response = second_response.to_json().unwrap(); - assert!( - server - .try_dispatch_jsonrpc_message(second_response) - .is_none() - ); - assert_eq!( - second_rx.next().await.unwrap().to_json().unwrap(), - expected_second_response - ); - }); + fn accepts_only_both_media_types() { + let mut headers = HeaderMap::new(); + headers.insert( + "accept", + "application/json, text/event-stream".parse().unwrap(), + ); + assert!(accepts_both(&headers)); + headers.insert("accept", "application/json".parse().unwrap()); + assert!(!accepts_both(&headers)); + headers.append("accept", "text/event-stream;q=0.9".parse().unwrap()); + assert!(accepts_both(&headers)); + headers.insert( + "accept", + "application/json, text/event-stream;q=0".parse().unwrap(), + ); + assert!(!accepts_both(&headers)); } #[test] - fn late_response_to_disconnected_post_cannot_reach_reused_id() { - futures::executor::block_on(async { - fn request(method: &str) -> RpcRequest { - let RawJsonRpcMessage::Request(request) = RawJsonRpcMessage::request( - method.to_owned(), - serde_json::json!({}), - RequestId::Null, - ) - .unwrap() else { - unreachable!("request constructor always returns a request") - }; - request - } - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - let (first_tx, first_rx) = mpsc::unbounded(); - let (second_tx, mut second_rx) = mpsc::unbounded(); - - server - .handle_http_message( - HttpMessage::Request { - http_request_id: uuid::Uuid::new_v4(), - request: request("example/first"), - response_tx: first_tx, - }, - &mut channel_tx, - ) - .unwrap(); - server - .handle_http_message( - HttpMessage::Request { - http_request_id: uuid::Uuid::new_v4(), - request: request("example/second"), - response_tx: second_tx, - }, - &mut channel_tx, - ) - .unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) - if request.method.as_ref() == "example/first" - )); - drop(first_rx); - - let callback = TransportFrame::Single( - RawJsonRpcMessage::notification("callback".into(), serde_json::json!({})).unwrap(), - ); - assert!(server.try_dispatch_jsonrpc_message(callback).is_none()); - assert!(matches!( - second_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Notification(_))) - )); - - let first_response = TransportFrame::Single(RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "first" })), - )); - assert!( - server - .try_dispatch_jsonrpc_message(first_response) - .is_none() + fn accept_quality_uses_http_decimal_grammar() { + for quality in ["1", "1.", "1.000", "0.001", "0.5", "0.999"] { + let mut headers = HeaderMap::new(); + headers.insert( + "accept", + format!("application/json;q={quality}, text/event-stream") + .parse() + .unwrap(), ); - server.activate_pending_calls(&mut channel_tx).unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) - if request.method.as_ref() == "example/second" - )); - - let second_response = TransportFrame::Single(RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "second" })), - )); - assert!( - server - .try_dispatch_jsonrpc_message(second_response) - .is_none() + assert!(accepts_both(&headers), "{quality}"); + } + for quality in [ + "NaN", "inf", "-1", "+1", "1e0", "0.0001", "1.001", "2", "", ".5", "00.5", "0", + "0.000", "1;q=0.9", + ] { + let mut headers = HeaderMap::new(); + headers.insert( + "accept", + format!("application/json;q={quality}, text/event-stream") + .parse() + .unwrap(), ); - assert!(matches!( - second_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Response( - RpcResponse::Result { result, .. } - ))) if result == serde_json::json!({ "source": "second" }) - )); - }); + assert!(!accepts_both(&headers), "{quality}"); + } } #[test] - fn forwards_batch_and_routes_grouped_response_without_flattening() { - futures::executor::block_on(async { - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - let (response_tx, mut response_rx) = mpsc::unbounded(); - let incoming = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([RawJsonRpcMessage::request( - "example".into(), - serde_json::json!({}), - RequestId::Number(7), - ) - .unwrap()]) - .unwrap(), - ); - - server - .handle_http_message( - HttpMessage::Frame { - http_request_id: uuid::Uuid::new_v4(), - frame: incoming, - request_ids: vec![RequestId::Number(7)], - response_tx: Some(response_tx), - }, - &mut channel_tx, - ) - .unwrap(); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Batch(_)) - )); - - let callback = TransportFrame::Single( - RawJsonRpcMessage::notification("callback".into(), serde_json::json!({})).unwrap(), - ); - assert!(server.try_dispatch_jsonrpc_message(callback).is_none()); - assert!(matches!( - response_rx.next().await, - Some(TransportFrame::Single(RawJsonRpcMessage::Notification(_))) - )); - - let frame = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::response( - RequestId::Number(7), - Ok(serde_json::json!({ "ok": true })), - ), - ]) - .unwrap(), - ); - let expected = frame.to_json().unwrap(); - - assert!(server.try_dispatch_jsonrpc_message(frame).is_none()); - let received = response_rx - .next() - .await - .expect("waiting HTTP request stays open"); - assert_eq!(received.to_json().unwrap(), expected); - assert!(matches!(received, TransportFrame::Batch(_))); - }); + fn mirrored_names_decode_canonical_base64() { + assert!(matches_mirror( + Some("=?base64?SGVsbG8sIOS4lueVjA==?="), + "Hello, 世界" + )); + assert!(matches_mirror(Some("simple"), "simple")); + assert!(!matches_mirror(Some("=?base64?SGVsbG8=?="), "different")); + assert!(!matches_mirror(Some("=?base64?SGVsbG8==?="), "Hello")); + assert!(!matches_mirror( + Some("=?base64?literal?="), + "=?base64?literal?=" + )); + assert!(matches_mirror( + Some("=?base64?unfinished"), + "=?base64?unfinished" + )); } #[test] - fn batch_post_round_trips_as_one_grouped_sse_response() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - let incoming = serde_json::json!([ - { - "jsonrpc": "2.0", - "id": 7, - "method": "example/first", - "params": {} - }, - { - "jsonrpc": "2.0", - "id": 8, - "method": "example/second", - "params": {} - } - ]); - - let http_response = handle_post(State(state), incoming.to_string()) - .await - .expect("batch POST should open an SSE response"); - let registration = registration_rx - .next() - .await - .expect("batch POST should register with the bridge"); - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - server - .handle_http_message(registration, &mut channel_tx) - .expect("batch POST should be forwarded to the channel"); - let forwarded = channel_rx - .next() - .await - .expect("channel should receive the batch frame"); - assert!(matches!(&forwarded, TransportFrame::Batch(_))); - assert_eq!( - serde_json::from_str::(&forwarded.to_json().unwrap()).unwrap(), - incoming - ); - - let response = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::response( - RequestId::Number(7), - Ok(serde_json::json!({ "source": "first" })), - ), - RawJsonRpcMessage::response( - RequestId::Number(8), - Ok(serde_json::json!({ "source": "second" })), - ), - ]) - .expect("grouped response should be non-empty"), - ); - assert!(server.try_dispatch_jsonrpc_message(response).is_none()); - - let payload = single_sse_payload(http_response).await; - let entries = payload - .as_array() - .expect("SSE payload should remain one JSON-RPC array"); - assert_eq!(entries.len(), 2); - assert_eq!(entries[0]["id"], 7); - assert_eq!(entries[0]["result"]["source"], "first"); - assert_eq!(entries[1]["id"], 8); - assert_eq!(entries[1]["result"]["source"], "second"); + fn response_preserves_mrtr_and_opaque_request_state() { + let result = serde_json::json!({ + "inputRequests": [{"method":"elicitation/create","params":{"message":"answer"}}], + "requestState": {"opaque": [1, 2, 3]}, + "_meta": {"trace":"preserve", "io.modelcontextprotocol/subscriptionId":"unrelated"}, + "subscriptionId": "internal-id" }); + let response = rpc_result(serde_json::json!(42), "internal-id", result.clone()); + assert_eq!(response["result"], result); + assert_eq!(response["id"], 42); + let mapped = rpc_result( + serde_json::json!("external"), + "internal-id", + serde_json::json!({"subscriptionId":"internal-id","requestState":"unchanged", + "_meta":{"trace":"preserve", "io.modelcontextprotocol/subscriptionId":"internal-id", + "progressToken":"internal-id"}}), + ); + assert_eq!(mapped["result"]["subscriptionId"], "internal-id"); + assert_eq!( + mapped["result"]["_meta"]["io.modelcontextprotocol/subscriptionId"], + "external" + ); + assert_eq!(mapped["result"]["_meta"]["progressToken"], "internal-id"); + assert_eq!(mapped["result"]["requestState"], "unchanged"); } #[test] - fn malformed_batch_cannot_steal_a_valid_null_id_batch_response() { - futures::executor::block_on(async { - let (registration_tx, mut registration_rx) = mpsc::unbounded(); - let state = Arc::new(BridgeState { registration_tx }); - - let valid_http_response = handle_post( - State(state.clone()), - serde_json::json!([{ - "jsonrpc": "2.0", - "id": null, - "method": "example", - "params": {} - }]) - .to_string(), - ) - .await - .expect("valid null-ID batch POST"); - let malformed_http_response = handle_post(State(state), "[17,false]".to_owned()) - .await - .expect("malformed batch receives its own JSON-RPC error array"); - - let valid_batch = registration_rx - .next() - .await - .expect("valid batch is forwarded"); - assert!( - registration_rx.try_recv().is_err(), - "malformed-only batch must not register a bridge waiter" - ); - - let mut server = RunningServer::new(); - let (mut channel_tx, mut channel_rx) = mpsc::unbounded(); - server - .handle_http_message(valid_batch, &mut channel_tx) - .expect("forward valid null-ID batch"); - assert!(matches!( - channel_rx.next().await, - Some(TransportFrame::Batch(_)) - )); + fn concurrent_logical_ids_with_same_external_id_stay_request_scoped() { + let mut first = + serde_json::json!({"_meta":{"io.modelcontextprotocol/subscriptionId":"one"}}); + let mut second = + serde_json::json!({"_meta":{"io.modelcontextprotocol/subscriptionId":"two"}}); + rewrite_subscription_id(&mut first, "one", &serde_json::json!(7)); + rewrite_subscription_id(&mut second, "two", &serde_json::json!(7)); + assert_eq!(first["_meta"]["io.modelcontextprotocol/subscriptionId"], 7); + assert_eq!(second["_meta"]["io.modelcontextprotocol/subscriptionId"], 7); + let mut mismatch = + serde_json::json!({"_meta":{"io.modelcontextprotocol/subscriptionId":"two"}}); + rewrite_subscription_id(&mut mismatch, "one", &serde_json::json!("7")); + assert_eq!( + mismatch["_meta"]["io.modelcontextprotocol/subscriptionId"], + "two" + ); + } - let valid_response = TransportFrame::Batch( - agent_client_protocol::TransportBatch::from_messages([ - RawJsonRpcMessage::response( - RequestId::Null, - Ok(serde_json::json!({ "source": "valid" })), - ), - ]) - .expect("valid response batch is non-empty"), - ); - assert!( - server - .try_dispatch_jsonrpc_message(valid_response) - .is_none() + #[tokio::test] + async fn rejects_legacy_methods_and_invalid_headers_over_real_http() { + async fn exchange( + address: std::net::SocketAddr, + route: &str, + method: &str, + headers: &str, + body: &str, + ) -> String { + let mut stream = tokio::net::TcpStream::connect(address).await.unwrap(); + let request = format!( + "{method} /{route} HTTP/1.1\r\nHost: {address}\r\nConnection: close\r\n{headers}Content-Length: {}\r\n\r\n{body}", + body.len() ); - - let valid_payload = single_sse_payload(valid_http_response).await; - let valid_entries = valid_payload - .as_array() - .expect("valid response should remain a batch"); - assert_eq!(valid_entries.len(), 1); - assert_eq!(valid_entries[0]["id"], serde_json::Value::Null); - assert_eq!(valid_entries[0]["result"]["source"], "valid"); - - let malformed_payload = single_sse_payload(malformed_http_response).await; - let malformed_entries = malformed_payload - .as_array() - .expect("malformed response should be an error batch"); - assert_eq!(malformed_entries.len(), 2); - for entry in malformed_entries { - assert_eq!(entry["id"], serde_json::Value::Null); - assert_eq!( - entry["error"]["code"], - i32::from(agent_client_protocol::ErrorCode::InvalidRequest) - ); - } - }); + stream.write_all(request.as_bytes()).await.unwrap(); + let mut response = String::new(); + stream.read_to_string(&mut response).await.unwrap(); + response + } + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let (tx, _rx) = mpsc::channel(8); + let state = BridgeState::new(tx); + let (url, token) = state.declaration_url(address.port(), "server"); + let route = url.rsplit('/').next().unwrap(); + let task = tokio::spawn(run_http_listener(listener, state)); + let auth = format!("Authorization: Bearer {token}\r\n"); + let legacy = exchange(address, route, "GET", &auth, "").await; + assert!(legacy.starts_with("HTTP/1.1 405"), "{legacy}"); + let delete = exchange(address, route, "DELETE", &auth, "").await; + assert!(delete.starts_with("HTTP/1.1 405"), "{delete}"); + let invalid_origin = + exchange(address, route, "POST", "Origin: http://evil.test\r\n", "{}").await; + assert!( + invalid_origin.starts_with("HTTP/1.1 403"), + "{invalid_origin}" + ); + let invalid_get_origin = + exchange(address, route, "GET", "Origin: http://evil.test\r\n", "").await; + assert!( + invalid_get_origin.starts_with("HTTP/1.1 403"), + "{invalid_get_origin}" + ); + let invalid_auth = exchange(address, route, "POST", "", "{}").await; + assert!(invalid_auth.starts_with("HTTP/1.1 401"), "{invalid_auth}"); + assert!( + invalid_auth + .to_ascii_lowercase() + .contains("www-authenticate: bearer"), + "{invalid_auth}" + ); + let body = serde_json::json!({"jsonrpc":"2.0","id":1,"method":"tools/list", + "params":{"_meta":{"io.modelcontextprotocol/protocolVersion":VERSION, + "io.modelcontextprotocol/clientCapabilities":{}}}}) + .to_string(); + let headers = format!( + "{auth}Accept: application/json, text/event-stream\r\nContent-Type: application/json\r\nMCP-Protocol-Version: 2026-07-28\r\nMcp-Method: wrong/method\r\n" + ); + let mismatch = exchange(address, route, "POST", &headers, &body).await; + assert!(mismatch.starts_with("HTTP/1.1 400"), "{mismatch}"); + assert!(mismatch.contains("-32020"), "{mismatch}"); + let batch = exchange(address, route, "POST", &headers, "[]").await; + assert!(batch.starts_with("HTTP/1.1 400"), "{batch}"); + let headers = headers.replace("wrong/method", "tools/list"); + let fractional_id = body.replace("\"id\":1", "\"id\":1.5"); + let fractional = tokio::time::timeout( + std::time::Duration::from_secs(3), + exchange(address, route, "POST", &headers, &fractional_id), + ) + .await + .expect("an invalid request ID must be rejected before forwarding"); + assert!(fractional.starts_with("HTTP/1.1 400"), "{fractional}"); + assert!(fractional.contains("-32600"), "{fractional}"); + let error: Value = + serde_json::from_str(fractional.split("\r\n\r\n").nth(1).unwrap()).unwrap(); + assert!(error.get("id").is_none()); + task.abort(); } } diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs index feeb6624..26689737 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/mod.rs @@ -1,121 +1,117 @@ -//! MCP-over-ACP compatibility proxy. +//! Request-scoped MCP 2026-07-28 Streamable HTTP adapter for native ACP MCP servers. //! -//! This proxy adapts schema-native `McpServer::Acp` declarations for agents that do not -//! support the ACP MCP transport. It replaces those declarations with loopback HTTP bridges and -//! relays `mcp/connect`, `mcp/message`, and `mcp/disconnect` over ACP. -//! -//! Stable protocol v1 is supported by default. Enable the crate's -//! `unstable_protocol_v2` feature to use the same proxy in a draft-v2 conductor -//! chain. -//! -//! # Usage -//! -//! ```rust,ignore -//! use agent_client_protocol_polyfill::mcp_over_acp::McpOverAcpPolyfill; -//! -//! let conductor = ConductorImpl::new_agent( -//! "conductor", -//! ProxiesAndAgent::new(my_agent).proxy(McpOverAcpPolyfill::http()), -//! ); -//! ``` +//! Native-capable successors receive the original declarations and messages unchanged. +//! HTTP-only successors receive loopback endpoints; no MCP connection or session is created. -mod actor; pub(crate) mod http; mod protocol; -use std::collections::HashMap; +use std::{ + collections::HashMap, + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, +}; use agent_client_protocol::{ Agent, Client, Conductor, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - Proxy, Responder, UntypedMessage, is_cancel_request_notification, util::MatchDispatchFrom, + Proxy, UntypedMessage, util::MatchDispatchFrom, +}; +use futures::{ + SinkExt, StreamExt, + channel::{mpsc, oneshot}, }; -use futures::{SinkExt, channel::mpsc, channel::oneshot}; use serde_json::Value; -use tokio::net::TcpListener; -use tracing::{debug, info, warn}; +use tokio::{net::TcpListener, sync::mpsc as tokio_mpsc}; +use tracing::{debug, warn}; -use self::actor::BridgeConnectionActor; use self::protocol::{ - DownstreamMcpMode, NativeMcpMessage, NativeServer, PolyfillProtocol, native_params_into_value, + DownstreamMcpMode, NativeMcpNotification, NativeMcpOutcome, PolyfillProtocol, }; -/// Internal messages for the polyfill's bridge management. -#[derive(Debug)] -pub(crate) enum BridgeMessage { - /// Record the selected ACP schema and which MCP transport the successor can consume. +// Conservative per-bridge limits. Notifications are bounded per HTTP POST by +// both message count and serialized bytes; terminal responses bypass the queue. +const MAX_ACTIVE_REQUESTS: usize = 64; +const MAX_QUEUED_NOTIFICATIONS: usize = 16; +const MAX_QUEUED_BYTES: usize = 256 * 1024; +const MAX_TERMINAL_BYTES: usize = 1024 * 1024; +const LOCAL_LIMIT_ERROR: i64 = -33000; + +struct QueuedNotification { + value: Value, + bytes: usize, + used: Arc, +} + +impl Drop for QueuedNotification { + fn drop(&mut self) { + self.used.fetch_sub(self.bytes, Ordering::Relaxed); + } +} + +#[derive(Clone)] +struct StreamSender { + tx: tokio_mpsc::Sender, + used: Arc, +} + +impl StreamSender { + fn send(&self, value: Value) -> Result<(), ()> { + let bytes = serde_json::to_vec(&value).map_err(|_| ())?.len(); + let reserved = self + .used + .fetch_update(Ordering::Relaxed, Ordering::Relaxed, |used| { + used.checked_add(bytes) + .filter(|total| *total <= MAX_QUEUED_BYTES) + }); + if reserved.is_err() { + return Err(()); + } + self.tx + .try_send(QueuedNotification { + value, + bytes, + used: self.used.clone(), + }) + .map_err(|_| ()) + } + + async fn closed(&self) { + self.tx.closed().await; + } +} + +enum BridgeMessage { SetProtocol { protocol: PolyfillProtocol, downstream_mode: DownstreamMcpMode, }, - - /// Transform the MCP declarations for one session setup request. TransformServers { servers: Vec, response_tx: oneshot::Sender, agent_client_protocol::Error>>, }, - - /// A new TCP connection was accepted and needs a native MCP connection ID. - ConnectionReceived { + Request { server_id: String, - actor: BridgeConnectionActor, - connection: BridgeConnection, + request_id: String, + http_id: Value, + method: String, + params: Option>, + response_tx: StreamSender, + terminal_tx: tokio::sync::oneshot::Sender, }, - - /// A native MCP connection ID was received; spawn the actor and store its sender. - ConnectionEstablished { - server_id: String, - connection_id: String, - actor: BridgeConnectionActor, - connection: BridgeConnection, + Notification(NativeMcpNotification), + Finished { + request_id: String, + result: Option>, }, - - /// Opening a native MCP connection failed. - ConnectionFailed { server_id: String }, - - /// An MCP message from the local agent that must be sent over ACP. - ClientToServer { - connection_id: String, - message: Dispatch, - }, - - /// An MCP server request received over ACP for the local agent's MCP client. - ServerToClientRequest { - request: NativeMcpMessage, - responder: Responder, - }, - - /// An MCP server notification received over ACP for the local agent's MCP client. - ServerToClientNotification { notification: NativeMcpMessage }, - - /// The local MCP bridge disconnected. - Disconnected { connection_id: String }, -} - -/// Connection handle for sending messages to an MCP client via a bridge. -#[derive(Clone, Debug)] -pub(crate) struct BridgeConnection { - to_mcp_client_tx: mpsc::Sender, -} - -impl BridgeConnection { - pub fn new(to_mcp_client_tx: mpsc::Sender) -> Self { - Self { to_mcp_client_tx } - } - - fn try_send(&mut self, message: Dispatch) -> Option> { - self.to_mcp_client_tx - .try_send(message) - .err() - .map(|error| Box::new(error.into_inner())) - } } -/// Adapts schema-native MCP-over-ACP declarations for agents that support HTTP MCP. +/// Adapts native MCP-over-ACP servers to loopback Streamable HTTP for HTTP-only agents. #[derive(Debug, Default)] pub struct McpOverAcpPolyfill; impl McpOverAcpPolyfill { - /// Create a polyfill that exposes each ACP MCP server through loopback HTTP. #[must_use] pub fn http() -> Self { Self @@ -136,7 +132,6 @@ impl ConnectTo for McpOverAcpPolyfill { .connect_to(client) .await } - #[cfg(not(feature = "unstable_protocol_v2"))] { McpOverAcpProxy(PolyfillProtocol::V1) @@ -160,14 +155,13 @@ impl ConnectTo for McpOverAcpProxy { bridge_rx, protocol: None, downstream_mode: DownstreamMcpMode::Unknown, - listeners: BridgeListeners::default(), - bridge_connections: HashMap::new(), + listener: None, + active: HashMap::new(), }; let handler = PolyfillHandler { protocol: None, bridge_tx, }; - match self.0 { PolyfillProtocol::V1 => { Proxy @@ -224,11 +218,81 @@ impl PolyfillHandler { cx: &ConnectionTo, ) -> Result, agent_client_protocol::Error> { match message { - Dispatch::Request(request, responder) => { - self.handle_client_request(request, responder, cx).await + Dispatch::Request(mut request, responder) => { + if request.method() == agent_client_protocol::schema::METHOD_INITIALIZE_PROXY { + if self.protocol.is_some() { + return Err(agent_client_protocol::Error::invalid_request() + .data("MCP-over-ACP polyfill was already initialized")); + } + let protocol = PolyfillProtocol::from_initialize_request(&request)?; + self.protocol = Some(protocol); + request.method = "initialize".into(); + let sent = cx + .send_request_to(Agent, request) + .forward_cancellation_from(responder.cancellation()); + let mut bridge_tx = self.bridge_tx.clone(); + sent.on_receiving_result(async move |result| { + let result = match result { + Ok(mut response) => { + let mode = protocol.transform_initialize_response(&mut response)?; + bridge_tx + .send(BridgeMessage::SetProtocol { + protocol, + downstream_mode: mode, + }) + .await + .map_err(agent_client_protocol::Error::into_internal_error)?; + Ok(response) + } + Err(error) => Err(error), + }; + responder.respond_with_result(result) + })?; + return Ok(Handled::Yes); + } + let Some(protocol) = self.protocol else { + return Ok(Handled::No { + message: Dispatch::Request(request, responder), + retry: false, + }); + }; + if protocol.is_session_setup_method(request.method()) { + protocol.validate_session_setup_request(&request)?; + transform_session_servers(&mut request, &mut self.bridge_tx).await?; + cx.send_request_to(Agent, request) + .forward_response_to(responder)?; + return Ok(Handled::Yes); + } + // Only agent-to-provider requests are valid; reverse RPC is never forwarded. + if request.method() == "mcp/message" { + responder + .respond_with_error(agent_client_protocol::Error::method_not_found())?; + return Ok(Handled::Yes); + } + Ok(Handled::No { + message: Dispatch::Request(request, responder), + retry: false, + }) } Dispatch::Notification(notification) => { - self.handle_client_notification(notification).await + if notification.method() == "mcp/message" { + let Some(protocol) = self.protocol else { + return Ok(Handled::No { + message: Dispatch::Notification(notification), + retry: false, + }); + }; + let notification = protocol.parse_notification(notification)?; + self.bridge_tx + .send(BridgeMessage::Notification(notification)) + .await + .map_err(agent_client_protocol::Error::into_internal_error)?; + return Ok(Handled::Yes); + } + Ok(Handled::No { + message: Dispatch::Notification(notification), + retry: false, + }) } message @ Dispatch::Response(_, _) => Ok(Handled::No { message, @@ -236,108 +300,6 @@ impl PolyfillHandler { }), } } - - async fn handle_client_request( - &mut self, - mut request: UntypedMessage, - responder: Responder, - cx: &ConnectionTo, - ) -> Result, agent_client_protocol::Error> { - if request.method() == agent_client_protocol::schema::METHOD_INITIALIZE_PROXY { - if self.protocol.is_some() { - return Err(agent_client_protocol::Error::invalid_request() - .data("MCP-over-ACP polyfill was already initialized")); - } - let protocol = PolyfillProtocol::from_initialize_request(&request)?; - self.protocol = Some(protocol); - request.method = "initialize".to_string(); - - let sent = cx.send_request_to(Agent, request); - let sent = sent.forward_cancellation_from(responder.cancellation()); - let mut bridge_tx = self.bridge_tx.clone(); - sent.on_receiving_result(async move |result| { - let result = match result { - Ok(response) => { - adapt_initialize_response(protocol, response, &mut bridge_tx).await - } - Err(error) => Err(error), - }; - responder.respond_with_result(result) - })?; - return Ok(Handled::Yes); - } - - let Some(protocol) = self.protocol else { - return Ok(Handled::No { - message: Dispatch::Request(request, responder), - retry: false, - }); - }; - - if protocol.is_session_setup_method(request.method()) { - protocol.validate_session_setup_request(&request)?; - transform_session_servers(&mut request, &mut self.bridge_tx).await?; - cx.send_request_to(Agent, request) - .forward_response_to(responder)?; - return Ok(Handled::Yes); - } - - if request.method() == "mcp/message" { - let request = protocol.parse_message_request(request)?; - self.bridge_tx - .send(BridgeMessage::ServerToClientRequest { request, responder }) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - return Ok(Handled::Yes); - } - - Ok(Handled::No { - message: Dispatch::Request(request, responder), - retry: false, - }) - } - - async fn handle_client_notification( - &mut self, - notification: UntypedMessage, - ) -> Result, agent_client_protocol::Error> { - let Some(protocol) = self.protocol else { - return Ok(Handled::No { - message: Dispatch::Notification(notification), - retry: false, - }); - }; - - if notification.method() == "mcp/message" { - let notification = protocol.parse_message_notification(notification)?; - self.bridge_tx - .send(BridgeMessage::ServerToClientNotification { notification }) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - return Ok(Handled::Yes); - } - - Ok(Handled::No { - message: Dispatch::Notification(notification), - retry: false, - }) - } -} - -async fn adapt_initialize_response( - protocol: PolyfillProtocol, - mut response: Value, - bridge_tx: &mut mpsc::Sender, -) -> Result { - let downstream_mode = protocol.transform_initialize_response(&mut response)?; - bridge_tx - .send(BridgeMessage::SetProtocol { - protocol, - downstream_mode, - }) - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - Ok(response) } async fn transform_session_servers( @@ -352,7 +314,6 @@ async fn transform_session_servers( else { return Ok(()); }; - let (response_tx, response_rx) = oneshot::channel(); bridge_tx .send(BridgeMessage::TransformServers { @@ -367,101 +328,14 @@ async fn transform_session_servers( Ok(()) } -#[derive(Default, Debug)] -struct BridgeListeners { - listeners: HashMap, -} - -#[derive(Clone, Debug)] -struct BridgeListener { - tcp_port: u16, -} - -impl BridgeListener { - fn declaration( - &self, - protocol: PolyfillProtocol, - server: NativeServer, - ) -> Result { - server.http_declaration(protocol, format!("http://127.0.0.1:{}", self.tcp_port)) - } -} - -impl BridgeListeners { - async fn transform_servers( - &mut self, - connection: &ConnectionTo, - protocol: PolyfillProtocol, - servers: Vec, - bridge_tx: &mpsc::Sender, - ) -> Result, agent_client_protocol::Error> { - let mut transformed = Vec::with_capacity(servers.len()); - for server in servers { - transformed.push( - self.transform_server(connection, protocol, server, bridge_tx) - .await?, - ); - } - Ok(transformed) - } - - async fn transform_server( - &mut self, - connection: &ConnectionTo, - protocol: PolyfillProtocol, - server: Value, - bridge_tx: &mpsc::Sender, - ) -> Result { - let Some(native_server) = protocol.native_server(server.clone()) else { - return Ok(server); - }; - let server_id = native_server.server_id.clone(); - - info!( - server_name = %native_server.name, - server_id, - "detected native MCP-over-ACP server; creating compatibility bridge" - ); - - if let Some(listener) = self.listeners.get(&server_id) { - return listener.declaration(protocol, native_server); - } - - let tcp_listener = TcpListener::bind("127.0.0.1:0") - .await - .map_err(agent_client_protocol::Error::into_internal_error)?; - let tcp_port = tcp_listener - .local_addr() - .map_err(agent_client_protocol::Error::into_internal_error)? - .port(); - let listener = BridgeListener { tcp_port }; - - connection.spawn({ - let server_id = server_id.clone(); - let bridge_tx = bridge_tx.clone(); - async move { - info!( - server_id, - tcp_port, "accepting MCP compatibility connections" - ); - http::run_http_listener(tcp_listener, server_id, bridge_tx).await - } - })?; - - let declaration = listener.declaration(protocol, native_server)?; - self.listeners.insert(server_id, listener); - Ok(declaration) - } - - fn remove(&mut self, server_id: &str) { - self.listeners.remove(server_id); - } -} - -#[derive(Debug)] -struct ActiveBridgeConnection { +struct ActiveRequest { + protocol: PolyfillProtocol, server_id: String, - bridge: BridgeConnection, + http_id: Value, + method: String, + response_tx: StreamSender, + terminal_tx: tokio::sync::oneshot::Sender, + cancel_tx: tokio::sync::oneshot::Sender<()>, } struct BridgeRunner { @@ -469,8 +343,8 @@ struct BridgeRunner { bridge_rx: mpsc::Receiver, protocol: Option, downstream_mode: DownstreamMcpMode, - listeners: BridgeListeners, - bridge_connections: HashMap, + listener: Option<(u16, Arc)>, + active: HashMap, } impl std::fmt::Debug for BridgeRunner { @@ -478,8 +352,8 @@ impl std::fmt::Debug for BridgeRunner { f.debug_struct("BridgeRunner") .field("protocol", &self.protocol) .field("downstream_mode", &self.downstream_mode) - .field("listeners", &self.listeners.listeners.len()) - .field("bridge_connections", &self.bridge_connections.len()) + .field("listener", &self.listener.is_some()) + .field("active", &self.active.len()) .finish_non_exhaustive() } } @@ -489,8 +363,6 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { mut self, connection: ConnectionTo, ) -> Result<(), agent_client_protocol::Error> { - use futures::StreamExt; - while let Some(message) = self.bridge_rx.next().await { match message { BridgeMessage::SetProtocol { @@ -500,281 +372,138 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { self.protocol = Some(protocol); self.downstream_mode = downstream_mode; } - BridgeMessage::TransformServers { servers, response_tx, } => { - let result = match (self.protocol, self.downstream_mode) { - (Some(_), DownstreamMcpMode::Native) => Ok(servers), - (Some(protocol), DownstreamMcpMode::HttpAdapter) => { - self.listeners - .transform_servers(&connection, protocol, servers, &self.bridge_tx) - .await - } - (Some(protocol), DownstreamMcpMode::Unavailable) => reject_native_servers( - protocol, - servers, - "the downstream agent supports neither native nor HTTP MCP transport", - ), - (Some(protocol), DownstreamMcpMode::Unknown) => reject_native_servers( - protocol, - servers, - "MCP transport capabilities are unavailable before initialize", - ), - (None, _) => Err(agent_client_protocol::Error::invalid_request() - .data("MCP transport capabilities are unavailable before initialize")), - }; + let result = self.transform_servers(&connection, servers).await; drop(response_tx.send(result)); } - - BridgeMessage::ConnectionReceived { + BridgeMessage::Request { server_id, - actor, - connection: bridge, + request_id, + http_id, + method, + params, + response_tx, + terminal_tx, } => { - let Some(protocol) = self.protocol else { - warn!( - server_id, - "cannot open MCP bridge before ACP initialization" - ); - self.listeners.remove(&server_id); + let Some(protocol) = self + .protocol + .filter(|_| self.downstream_mode == DownstreamMcpMode::HttpAdapter) + else { + drop(terminal_tx.send(http::rpc_error( + http_id, + -33002, + "MCP adapter unavailable", + ))); continue; }; - let request = protocol.connect_request(server_id.clone())?; - let mut bridge_tx = self.bridge_tx.clone(); - let scheduled = connection - .send_request_to(Client, request) - .on_receiving_result(async move |result| { - let message = match result { - Ok(response) => match protocol.connect_response_id(response) { - Ok(connection_id) => BridgeMessage::ConnectionEstablished { - server_id, - connection_id, - actor, - connection: bridge, - }, - Err(error) => { - warn!(?error, "invalid response to mcp/connect"); - BridgeMessage::ConnectionFailed { server_id } - } - }, - Err(error) => { - warn!(?error, "mcp/connect failed"); - BridgeMessage::ConnectionFailed { server_id } - } - }; - drop(bridge_tx.send(message).await); - Ok(()) - }); - if let Err(error) = scheduled { - warn!(?error, "could not schedule mcp/connect response handling"); + if !self.can_admit_request() { + drop(terminal_tx.send(http::rpc_error( + http_id, + LOCAL_LIMIT_ERROR, + "Too many active MCP requests", + ))); + continue; } - } - - BridgeMessage::ConnectionEstablished { - server_id, - connection_id, - actor, - connection: bridge, - } => { - self.bridge_connections.insert( - connection_id.clone(), - ActiveBridgeConnection { server_id, bridge }, + let (cancel_tx, cancel_rx) = tokio::sync::oneshot::channel(); + self.active.insert( + request_id.clone(), + ActiveRequest { + protocol, + server_id: server_id.clone(), + http_id: http_id.clone(), + method: method.clone(), + response_tx: response_tx.clone(), + terminal_tx, + cancel_tx, + }, ); - connection.spawn(actor.run(connection_id))?; - } - - BridgeMessage::ConnectionFailed { server_id } => { - self.listeners.remove(&server_id); - } - - BridgeMessage::ClientToServer { - connection_id, - message, - } => { - let Some(protocol) = self.protocol else { - let rejection = match message { - Dispatch::Request(_, responder) => responder - .respond_with_internal_error( - "ACP protocol is unavailable before initialize", - ), - Dispatch::Notification(_) | Dispatch::Response(_, _) => Ok(()), + let mut tx = self.bridge_tx.clone(); + let cx = connection.clone(); + let request_id_for_task = request_id.clone(); + connection.spawn(async move { + // Dropping the HTTP response stream cancels precisely this ACP request. + let result = tokio::select! { + result = forward_http_request(cx, protocol, server_id, + request_id_for_task, method, params) => Some(result), + () = response_tx.closed() => None, + _ = cancel_rx => None, }; - if let Err(error) = rejection { - debug!(?error, "could not reject MCP request before initialize"); - } - continue; - }; - - match message { - Dispatch::Request(message, responder) => { - match protocol.message_request(connection_id, message) { - Ok(request) => { - let pending = connection.send_request_to(Client, request); - if let Err(error) = pending.forward_response_to(responder) { - warn!( - ?error, - "could not forward local MCP request response" - ); - } - } - Err(error) => { - if let Err(send_error) = responder.respond_with_error(error) { - debug!( - ?send_error, - "could not reject malformed MCP request" - ); - } - } - } - } - Dispatch::Notification(message) => { - match local_mcp_notification(protocol, connection_id, message) { - Ok(Some(notification)) => { - if let Err(error) = - connection.send_notification_to(Client, notification) - { - warn!(?error, "could not forward local MCP notification"); - } - } - Ok(None) => { - debug!( - "not tunneling hop-scoped MCP cancellation through mcp/message" - ); - } - Err(error) => { - warn!(?error, "could not forward local MCP notification"); - } - } - } - Dispatch::Response(result, router) => { - if let Err(error) = router.route_with_result(result) { - debug!(?error, "could not route MCP client response"); - } - } - } + tx.send(BridgeMessage::Finished { request_id, result }) + .await + .map_err(agent_client_protocol::Error::into_internal_error)?; + Ok(()) + })?; } - - BridgeMessage::ServerToClientRequest { request, responder } => { - match self.downstream_mode { - DownstreamMcpMode::Native => { - let pending = connection.send_request_to(Agent, request.raw); - if let Err(error) = pending.forward_response_to(responder) { - debug!(?error, "could not forward native MCP request"); - } - } - DownstreamMcpMode::HttpAdapter => { - let connection_id = request.connection_id; - let Some(active) = self.bridge_connections.get_mut(&connection_id) - else { - respond_unknown_connection(responder, &connection_id); - continue; - }; - let message = UntypedMessage { - method: request.method, - params: native_params_into_value(request.params), - }; - if let Some(message) = active - .bridge - .try_send(Dispatch::Request(message, responder)) - { - let Dispatch::Request(_, responder) = *message else { - unreachable!("the failed bridge message was a request") - }; - if let Err(send_error) = responder.respond_with_internal_error( - "the local MCP client is unavailable or backpressured", - ) { - debug!( - ?send_error, - "could not reject unavailable MCP connection" - ); - } - } - } - DownstreamMcpMode::Unknown | DownstreamMcpMode::Unavailable => { - if let Err(error) = - responder.respond_with_error( - agent_client_protocol::Error::method_not_found(), - ) - { - debug!(?error, "could not reject unsupported native MCP request"); - } - } - } - } - - BridgeMessage::ServerToClientNotification { notification } => { - match self.downstream_mode { - DownstreamMcpMode::Native => { - if let Err(error) = - connection.send_notification_to(Agent, notification.raw) - { - debug!(?error, "could not forward native MCP notification"); - } - } - DownstreamMcpMode::HttpAdapter => { - let connection_id = notification.connection_id; - let Some(active) = self.bridge_connections.get_mut(&connection_id) - else { - debug!( - connection_id, - "ignoring notification for unknown MCP connection" - ); - continue; - }; - let message = UntypedMessage { - method: notification.method, - params: native_params_into_value(notification.params), - }; - if active - .bridge - .try_send(Dispatch::Notification(message)) - .is_some() - { - debug!("discarding MCP notification for unavailable local client"); - } + BridgeMessage::Notification(notification) => { + if self.downstream_mode == DownstreamMcpMode::Native { + connection.send_notification_to(Agent, notification.raw)?; + } else if self.downstream_mode == DownstreamMcpMode::HttpAdapter { + let Some(active) = self.active.get(¬ification.request_id) else { + debug!("dropping notification for stale MCP request"); + continue; + }; + if active.server_id != notification.server_id { + warn!("dropping notification with mismatched MCP server"); + continue; } - DownstreamMcpMode::Unknown | DownstreamMcpMode::Unavailable => { - debug!("ignoring unsupported native MCP notification"); + let mut params = Value::Object(notification.params.unwrap_or_default()); + http::rewrite_subscription_id( + &mut params, + ¬ification.request_id, + &active.http_id, + ); + let message = serde_json::json!({ + "jsonrpc": "2.0", + "method": notification.method, + "params": params, + }); + if active.response_tx.send(message).is_err() { + // Stop only this request. Its final error goes through a + // separate control path that cannot be blocked by a full queue. + let active = self + .active + .remove(¬ification.request_id) + .expect("active request checked above"); + let _ = active.cancel_tx.send(()); + drop(active.terminal_tx.send(http::rpc_error( + active.http_id, + LOCAL_LIMIT_ERROR, + "MCP notification queue overflow", + ))); } } } - - BridgeMessage::Disconnected { connection_id } => { - let Some(active) = self.bridge_connections.remove(&connection_id) else { - debug!(connection_id, "local MCP connection was already removed"); + BridgeMessage::Finished { request_id, result } => { + let Some(active) = self.active.remove(&request_id) else { continue; }; - self.listeners.remove(&active.server_id); - - let Some(protocol) = self.protocol else { - debug!("could not disconnect MCP bridge before ACP initialization"); - continue; - }; - let request = protocol.disconnect_request(connection_id)?; - let scheduled = connection - .send_request_to(Client, request) - .on_receiving_result(async move |result| { - match result { - Ok(response) => { - if let Err(error) = - protocol.validate_disconnect_response(response) - { - warn!(?error, "invalid response to mcp/disconnect"); - } - } - Err(error) => { - debug!(?error, "mcp/disconnect failed"); - } - } - Ok(()) - }); - if let Err(error) = scheduled { - debug!( - ?error, - "could not schedule mcp/disconnect response handling" - ); + if let Some(result) = result { + let http_id = active.http_id.clone(); + let value = match result { + Ok(carrier) => project_mcp_carrier( + active.protocol, + active.http_id, + &request_id, + &active.method, + carrier, + ), + Err(error) => http::rpc_binding_error(active.http_id, error), + }; + let value = if serde_json::to_vec(&value) + .is_ok_and(|bytes| bytes.len() <= MAX_TERMINAL_BYTES) + { + value + } else { + http::rpc_error( + http_id, + LOCAL_LIMIT_ERROR, + "MCP terminal response too large", + ) + }; + drop(active.terminal_tx.send(value)); } } } @@ -783,286 +512,319 @@ impl agent_client_protocol::RunWithConnectionTo for BridgeRunner { } } -fn local_mcp_notification( +/// ACP success carries exactly one MCP outcome. An outer ACP failure is a +/// binding/runtime failure, not an MCP error carried in a successful response. +fn project_mcp_carrier( protocol: PolyfillProtocol, - connection_id: String, - message: UntypedMessage, -) -> Result, agent_client_protocol::Error> { - if is_cancel_request_notification(&message) { - return Ok(None); + http_id: Value, + request_id: &str, + method: &str, + carrier: Value, +) -> Value { + match protocol.message_response(carrier) { + Ok(NativeMcpOutcome::Result(mut result)) => { + if method == "tools/list" { + strip_header_annotations(&mut result); + } + http::rpc_result(http_id, request_id, result) + } + Ok(NativeMcpOutcome::Error(error)) => http::rpc_peer_error(http_id, error), + _ => http::rpc_error(http_id, -33002, "Invalid MCP-over-ACP response carrier"), } - protocol - .message_notification(connection_id, message) - .map(Some) } -fn reject_native_servers( - protocol: PolyfillProtocol, - servers: Vec, - reason: &'static str, -) -> Result, agent_client_protocol::Error> { - if servers - .iter() - .any(|server| protocol.native_server(server.clone()).is_some()) - { - Err(agent_client_protocol::Error::invalid_params().data(reason)) - } else { - Ok(servers) +impl BridgeRunner { + fn can_admit_request(&self) -> bool { + self.active.len() < MAX_ACTIVE_REQUESTS } -} -fn respond_unknown_connection(responder: Responder, connection_id: &str) { - let error = agent_client_protocol::Error::invalid_params().data(serde_json::json!({ - "reason": "unknown MCP connection", - "connectionId": connection_id, - })); - if let Err(send_error) = responder.respond_with_error(error) { - debug!( - ?send_error, - connection_id, "could not reject unknown MCP connection" - ); + async fn transform_servers( + &mut self, + connection: &ConnectionTo, + servers: Vec, + ) -> Result, agent_client_protocol::Error> { + let protocol = self + .protocol + .ok_or_else(agent_client_protocol::Error::invalid_request)?; + let mut transformed = Vec::with_capacity(servers.len()); + for server in servers { + let Some(native) = protocol.native_server(server.clone()) else { + transformed.push(server); + continue; + }; + match self.downstream_mode { + DownstreamMcpMode::Native => transformed.push(server), + DownstreamMcpMode::HttpAdapter => { + if self.listener.is_none() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .map_err(agent_client_protocol::Error::into_internal_error)?; + let port = listener + .local_addr() + .map_err(agent_client_protocol::Error::into_internal_error)? + .port(); + let state = http::BridgeState::new(self.bridge_tx.clone()); + connection.spawn(http::run_http_listener(listener, state.clone()))?; + self.listener = Some((port, state)); + } + let (port, state) = self.listener.as_ref().expect("listener created"); + let (url, token) = state.declaration_url(*port, &native.server_id); + transformed.push(native.http_declaration(protocol, url, &token)?); + } + DownstreamMcpMode::Unknown | DownstreamMcpMode::Unavailable => { + return Err(agent_client_protocol::Error::invalid_params().data( + "the downstream agent supports neither native nor HTTP MCP transport", + )); + } + } + } + Ok(transformed) } } -#[cfg(test)] -mod tests { - use std::collections::HashMap; +async fn forward_http_request( + connection: ConnectionTo, + protocol: PolyfillProtocol, + server_id: String, + request_id: String, + method: String, + params: Option>, +) -> Result { + let request = protocol.message_request(server_id, request_id, method, params, None)?; + connection + .send_request_to(Client, request) + .block_task() + .await +} - use agent_client_protocol::{ - Conductor, Dispatch, ErrorCode, Proxy, UntypedMessage, - schema::v1::{ - McpServer, McpServerAcp, McpServerHttp, MessageMcpNotification, MessageMcpRequest, - }, +fn strip_header_annotations(result: &mut Value) { + let Some(tools) = result.get_mut("tools").and_then(Value::as_array_mut) else { + return; }; - use futures::{StreamExt, channel::mpsc}; + for tool in tools { + if let Some(schema) = tool.get_mut("inputSchema") { + strip_schema_annotation(schema); + } + } +} - use super::{ - ActiveBridgeConnection, BridgeConnection, BridgeListener, BridgeListeners, BridgeRunner, - DownstreamMcpMode, PolyfillHandler, PolyfillProtocol, local_mcp_notification, - reject_native_servers, +fn strip_schema_annotation(schema: &mut Value) { + let Some(object) = schema.as_object_mut() else { + return; }; + object.remove("x-mcp-header"); + for key in [ + "properties", + "patternProperties", + "$defs", + "definitions", + "dependentSchemas", + ] { + if let Some(children) = object.get_mut(key).and_then(Value::as_object_mut) { + for child in children.values_mut() { + strip_schema_annotation(child); + } + } + } + for key in [ + "items", + "additionalItems", + "additionalProperties", + "unevaluatedItems", + "unevaluatedProperties", + "contains", + "contentSchema", + "not", + "if", + "then", + "else", + "propertyNames", + ] { + if let Some(child) = object.get_mut(key) { + strip_schema_annotation(child); + } + } + for key in ["allOf", "anyOf", "oneOf", "prefixItems"] { + if let Some(children) = object.get_mut(key).and_then(Value::as_array_mut) { + for child in children { + strip_schema_annotation(child); + } + } + } +} - #[test] - fn http_declarations_reuse_endpoint_but_preserve_name_and_meta() { - let listener = BridgeListener { tcp_port: 4321 }; - let first_meta = serde_json::Map::from_iter([("source".into(), "first".into())]); - let second_meta = serde_json::Map::from_iter([("source".into(), "second".into())]); - - let first = PolyfillProtocol::V1 - .native_server( - serde_json::to_value(McpServer::Acp( - McpServerAcp::new("first", "shared").meta(first_meta.clone()), - )) - .unwrap(), - ) - .unwrap(); - let second = PolyfillProtocol::V1 - .native_server( - serde_json::to_value(McpServer::Acp( - McpServerAcp::new("second", "shared").meta(second_meta.clone()), - )) - .unwrap(), - ) - .unwrap(); - let first: McpServer = - serde_json::from_value(listener.declaration(PolyfillProtocol::V1, first).unwrap()) - .unwrap(); - let second: McpServer = - serde_json::from_value(listener.declaration(PolyfillProtocol::V1, second).unwrap()) - .unwrap(); +#[cfg(test)] +mod http_limits_tests { + use super::*; - let McpServer::Http(first) = first else { - panic!("expected HTTP declaration") - }; - let McpServer::Http(second) = second else { - panic!("expected HTTP declaration") - }; - assert_eq!(first.url, "http://127.0.0.1:4321"); - assert_eq!(second.url, first.url); - assert_eq!(first.name, "first"); - assert_eq!(second.name, "second"); - assert_eq!(first.meta, Some(first_meta)); - assert_eq!(second.meta, Some(second_meta)); + #[test] + fn annotation_removal_only_traverses_schema_locations() { + let mut listing = serde_json::json!({"tools":[{ + "name":"with-header", + "inputSchema":{ + "type":"object", + "properties":{ + "x-mcp-header":{"type":"string","default":"retain"}, + "nested":{"type":"object","x-mcp-header":"Nested","properties":{ + "value":{"type":"string","x-mcp-header":"Value", + "examples":[{"x-mcp-header":"user data"}]} + }} + }, + "$defs":{"inner":{"type":"string","x-mcp-header":"Inner"}}, + "default":{"x-mcp-header":"not a schema"} + } + }]}); + strip_header_annotations(&mut listing); + let schema = &listing["tools"][0]["inputSchema"]; + assert_eq!(schema["properties"]["x-mcp-header"]["default"], "retain"); + assert_eq!( + schema["properties"]["nested"]["properties"]["value"]["examples"][0]["x-mcp-header"], + "user data" + ); + assert_eq!(schema["default"]["x-mcp-header"], "not a schema"); + assert!(schema["properties"]["nested"].get("x-mcp-header").is_none()); + assert!(schema["$defs"]["inner"].get("x-mcp-header").is_none()); } #[test] - fn downstream_mode_prefers_native_then_http_adaptation() { - assert_eq!( - DownstreamMcpMode::from_capabilities(true, true), - DownstreamMcpMode::Native + fn slow_reader_overflows_by_count_without_blocking_other_requests() { + let (tx, mut rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); + let sender = StreamSender { + tx, + used: Arc::new(AtomicUsize::new(0)), + }; + let (other_tx, mut other_rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); + let other = StreamSender { + tx: other_tx, + used: Arc::new(AtomicUsize::new(0)), + }; + for i in 0..MAX_QUEUED_NOTIFICATIONS { + assert!(sender.send(serde_json::json!({"sequence":i})).is_ok()); + } + assert!( + sender + .send(serde_json::json!({"sequence":"overflow"})) + .is_err() ); - assert_eq!( - DownstreamMcpMode::from_capabilities(false, true), - DownstreamMcpMode::Native + assert!( + other + .send(serde_json::json!({"sequence":"unaffected"})) + .is_ok() ); - assert_eq!( - DownstreamMcpMode::from_capabilities(true, false), - DownstreamMcpMode::HttpAdapter + assert_eq!(other_rx.try_recv().unwrap().value["sequence"], "unaffected"); + while rx.try_recv().is_ok() {} + assert_eq!(sender.used.load(Ordering::Relaxed), 0); + assert!( + sender + .send(serde_json::json!({"sequence":"recovered"})) + .is_ok() ); - assert_eq!( - DownstreamMcpMode::from_capabilities(false, false), - DownstreamMcpMode::Unavailable + } + + #[test] + fn large_notification_exceeds_byte_budget_without_reserving_memory() { + let (tx, _rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); + let sender = StreamSender { + tx, + used: Arc::new(AtomicUsize::new(0)), + }; + assert!( + sender + .send(serde_json::json!({"data":"x".repeat(MAX_QUEUED_BYTES)})) + .is_err() ); + assert_eq!(sender.used.load(Ordering::Relaxed), 0); + assert!(sender.send(serde_json::json!({"data":"ok"})).is_ok()); } #[test] - fn local_cancellation_is_not_tunneled_as_an_mcp_message() { - let cancellation = UntypedMessage { - method: "$/cancel_request".to_string(), - params: serde_json::json!({ - "requestId": "loopback-request" - }), + fn backend_capacity_reopens_when_an_active_request_finishes() { + let (bridge_tx, bridge_rx) = mpsc::channel(1); + let mut runner = BridgeRunner { + bridge_tx, + bridge_rx, + protocol: None, + downstream_mode: DownstreamMcpMode::Unknown, + listener: None, + active: HashMap::new(), }; - assert_eq!( - local_mcp_notification( + let (tx, _rx) = tokio_mpsc::channel(MAX_QUEUED_NOTIFICATIONS); + let sender = StreamSender { + tx, + used: Arc::new(AtomicUsize::new(0)), + }; + for index in 0..MAX_ACTIVE_REQUESTS { + let (terminal_tx, _terminal_rx) = tokio::sync::oneshot::channel(); + let (cancel_tx, _cancel_rx) = tokio::sync::oneshot::channel(); + runner.active.insert( + index.to_string(), + ActiveRequest { + protocol: PolyfillProtocol::V1, + server_id: String::new(), + http_id: Value::Null, + method: String::new(), + response_tx: sender.clone(), + terminal_tx, + cancel_tx, + }, + ); + } + assert!(!runner.can_admit_request()); + runner.active.remove("0"); + assert!(runner.can_admit_request()); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn mcp_carrier_preserves_peer_error_and_rejects_ambiguous_outcomes() { + let error = serde_json::json!({ + "code":-32000,"message":"peer-defined error", + "data":{"nested":[1,2]},"extension":"preserved" + }); + let project = |carrier| { + project_mcp_carrier( PolyfillProtocol::V1, - "native-connection".to_string(), - cancellation, + serde_json::json!("external"), + "internal", + "tools/call", + carrier, ) - .expect("cancellation filtering should not fail"), - None - ); - - let notification = UntypedMessage { - method: "notifications/progress".to_string(), - params: serde_json::json!({ - "progressToken": "token", - "progress": 0.5 - }), }; - let wrapped = local_mcp_notification( - PolyfillProtocol::V1, - "native-connection".to_string(), - notification, - ) - .expect("the notification should serialize") - .expect("ordinary MCP notifications should be forwarded"); - assert_eq!(wrapped.method, "mcp/message"); - assert_eq!( - wrapped.params["connectionId"], - serde_json::json!("native-connection") - ); + assert_eq!(project(serde_json::json!({"error":error}))["error"], error); assert_eq!( - wrapped.params["method"], - serde_json::json!("notifications/progress") + project(serde_json::json!({"result":null})), + serde_json::json!({"jsonrpc":"2.0","id":"external","result":null}) ); + for invalid in [ + serde_json::json!({"result":null,"error":error}), + serde_json::json!({"tools":[]}), + serde_json::json!({"error":null}), + ] { + let response = project(invalid); + assert_eq!(response["id"], "external"); + assert_eq!(response["error"]["code"], -33002); + } } #[test] - fn unavailable_mode_rejects_only_native_declarations() { - let standard = vec![ - serde_json::to_value(McpServer::Http(McpServerHttp::new( - "remote", - "https://example.com/mcp", - ))) - .unwrap(), - ]; + fn annotated_tools_are_reexported_without_transport_annotations() { + let mut result = serde_json::json!({"tools":[ + {"name":"plain","inputSchema":{"type":"object","properties":{}}}, + {"name":"annotated","inputSchema":{"properties":{"nested":{"properties":{ + "region":{"type":"string","x-mcp-header":"Region"} + }}}}} + ]}); + strip_header_annotations(&mut result); + assert_eq!(result["tools"].as_array().unwrap().len(), 2); + assert_eq!(result["tools"][0]["name"], "plain"); + assert_eq!(result["tools"][1]["name"], "annotated"); assert_eq!( - reject_native_servers(PolyfillProtocol::V1, standard.clone(), "unsupported").unwrap(), - standard + result["tools"][1]["inputSchema"]["properties"]["nested"]["properties"]["region"], + serde_json::json!({"type":"string"}) ); - - let error = reject_native_servers( - PolyfillProtocol::V1, - vec![ - serde_json::to_value(McpServer::Acp(McpServerAcp::new("native", "server-1"))) - .unwrap(), - ], - "unsupported", - ) - .expect_err("native declarations require a downstream transport"); - assert_eq!(error.code, ErrorCode::InvalidParams); - assert_eq!(error.data, Some(serde_json::json!("unsupported"))); - } - - #[tokio::test(flavor = "current_thread")] - async fn reverse_messages_route_without_stopping_on_unknown_connections() - -> Result<(), agent_client_protocol::Error> { - let known_connection_id = "known-connection"; - let (bridge_tx, bridge_rx) = mpsc::channel(16); - let (to_mcp_client_tx, mut to_mcp_client_rx) = mpsc::channel(16); - let bridge_connections = HashMap::from([( - known_connection_id.to_string(), - ActiveBridgeConnection { - server_id: "test-server".to_string(), - bridge: BridgeConnection::new(to_mcp_client_tx), - }, - )]); - - let proxy = Proxy - .builder() - .with_runner(BridgeRunner { - bridge_tx: bridge_tx.clone(), - bridge_rx, - protocol: Some(PolyfillProtocol::V1), - downstream_mode: DownstreamMcpMode::HttpAdapter, - listeners: BridgeListeners::default(), - bridge_connections, - }) - .with_handler(PolyfillHandler { - protocol: Some(PolyfillProtocol::V1), - bridge_tx, - }); - - Conductor - .builder() - .connect_with(proxy, async move |connection| { - let request_params = serde_json::Map::from_iter([( - "cursor".to_string(), - serde_json::json!("next-page"), - )]); - let request = MessageMcpRequest::new(known_connection_id, "tools/list") - .params(request_params.clone()); - let pending_response = connection.send_request(request); - - let Some(Dispatch::Request(message, responder)) = to_mcp_client_rx.next().await - else { - panic!("expected the request to reach the stored bridge connection") - }; - assert_eq!(message.method, "tools/list"); - assert_eq!(message.params, serde_json::Value::Object(request_params)); - - let inner_response = serde_json::json!({"tools": [{"name": "echo"}]}); - responder.respond(inner_response.clone())?; - let response = pending_response.block_task().await?; - let response: serde_json::Value = serde_json::from_str(response.0.get())?; - assert_eq!(response, inner_response); - - let unknown_error = connection - .send_request(MessageMcpRequest::new( - "missing-connection", - "resources/list", - )) - .block_task() - .await - .expect_err("an unknown connection must receive an error response"); - assert_eq!(unknown_error.code, ErrorCode::InvalidParams); - assert_eq!( - unknown_error.data, - Some(serde_json::json!({ - "reason": "unknown MCP connection", - "connectionId": "missing-connection", - })) - ); - - connection.send_notification(MessageMcpNotification::new( - "missing-connection", - "notifications/progress", - ))?; - connection.send_notification(MessageMcpNotification::new( - known_connection_id, - "notifications/tools/list_changed", - ))?; - - let Some(Dispatch::Notification(notification)) = to_mcp_client_rx.next().await - else { - panic!("expected the known notification after ignoring the unknown one") - }; - assert_eq!(notification.method, "notifications/tools/list_changed"); - assert_eq!(notification.params, serde_json::Value::Null); - - Ok(()) - }) - .await } } diff --git a/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs b/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs index b3cedc1b..ea595f4a 100644 --- a/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs +++ b/src/agent-client-protocol-polyfill/src/mcp_over_acp/protocol.rs @@ -2,22 +2,16 @@ use agent_client_protocol::{ Error, JsonRpcMessage, JsonRpcResponse, UntypedMessage, schema::{ InitializeProxyRequest, METHOD_INITIALIZE_PROXY, ProtocolVersion, - v1::{ - ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, DisconnectMcpResponse, - LoadSessionRequest, McpServer, MessageMcpNotification, MessageMcpRequest, - NewSessionRequest, ResumeSessionRequest, - }, + v1::{self, LoadSessionRequest, McpServer, NewSessionRequest, ResumeSessionRequest}, }, }; use serde_json::{Map, Value}; -#[cfg(feature = "unstable_protocol_v2")] -use agent_client_protocol::schema::v2; - #[cfg(feature = "unstable_session_fork")] use agent_client_protocol::schema::v1::ForkSessionRequest; +#[cfg(feature = "unstable_protocol_v2")] +use agent_client_protocol::schema::v2; -/// ACP schema selected by the conductor's proxy initialization request. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(crate) enum PolyfillProtocol { V1, @@ -25,23 +19,51 @@ pub(crate) enum PolyfillProtocol { V2, } +pub(super) enum NativeMcpOutcome { + Result(Value), + Error(Value), +} + impl PolyfillProtocol { + /// Validate against the negotiated ACP version before projecting onto HTTP. + pub(super) fn message_response(self, value: Value) -> Result { + match self { + Self::V1 => match v1::MessageMcpResponse::from_value("mcp/message", value)? { + v1::MessageMcpResponse::Result { result, .. } => { + Ok(NativeMcpOutcome::Result(result)) + } + v1::MessageMcpResponse::Error { error, .. } => { + Ok(NativeMcpOutcome::Error(serde_json::to_value(error)?)) + } + _ => Err(Error::invalid_request().data("unsupported MCP outcome")), + }, + #[cfg(feature = "unstable_protocol_v2")] + Self::V2 => match v2::MessageMcpResponse::from_value("mcp/message", value)? { + v2::MessageMcpResponse::Result { result, .. } => { + Ok(NativeMcpOutcome::Result(result)) + } + v2::MessageMcpResponse::Error { error, .. } => { + Ok(NativeMcpOutcome::Error(serde_json::to_value(error)?)) + } + _ => Err(Error::invalid_request().data("unsupported MCP outcome")), + }, + } + } + pub(crate) fn from_initialize_request(request: &UntypedMessage) -> Result { if request.method() != METHOD_INITIALIZE_PROXY { - return Err(Error::invalid_request() - .data(format!("expected `{METHOD_INITIALIZE_PROXY}` request"))); + return Err(Error::invalid_request().data("expected initialize proxy request")); } - - let requested = request - .params() - .get("protocolVersion") - .cloned() - .ok_or_else(invalid_initialize_protocol_version) - .and_then(|version| { - serde_json::from_value::(version) - .map_err(|_| invalid_initialize_protocol_version()) - })?; - + let requested = serde_json::from_value::( + request + .params() + .get("protocolVersion") + .cloned() + .ok_or_else(|| { + Error::invalid_params().data("missing initialize.protocolVersion") + })?, + ) + .map_err(Error::into_internal_error)?; let protocol = if requested == ProtocolVersion::V1 { Self::V1 } else { @@ -50,22 +72,17 @@ impl PolyfillProtocol { if requested == ProtocolVersion::V2 { Self::V2 } else { - return Err(unsupported_protocol_version(requested)); + return Err(Error::invalid_request() + .data(format!("unsupported ACP protocol version {requested}"))); } } - #[cfg(not(feature = "unstable_protocol_v2"))] { - return Err(unsupported_protocol_version(requested)); + return Err(Error::invalid_request() + .data(format!("unsupported ACP protocol version {requested}"))); } }; - - protocol.validate_initialize_request(request)?; - Ok(protocol) - } - - fn validate_initialize_request(self, request: &UntypedMessage) -> Result<(), Error> { - match self { + match protocol { Self::V1 => { InitializeProxyRequest::parse_message(request.method(), request.params())?; } @@ -74,7 +91,7 @@ impl PolyfillProtocol { v2::InitializeProxyRequest::parse_message(request.method(), request.params())?; } } - Ok(()) + Ok(protocol) } pub(crate) fn transform_initialize_response( @@ -83,72 +100,56 @@ impl PolyfillProtocol { ) -> Result { let mode = match self { Self::V1 => { - let response = agent_client_protocol::schema::v1::InitializeResponse::from_value( + let parsed = agent_client_protocol::schema::v1::InitializeResponse::from_value( "initialize", response.clone(), )?; DownstreamMcpMode::from_capabilities( - response.agent_capabilities.mcp_capabilities.http, - response.agent_capabilities.mcp_capabilities.acp, + parsed.agent_capabilities.mcp_capabilities.http, + parsed.agent_capabilities.mcp_capabilities.acp, ) } #[cfg(feature = "unstable_protocol_v2")] Self::V2 => { - let response = v2::InitializeResponse::from_value("initialize", response.clone())?; - let mcp = response + let parsed = v2::InitializeResponse::from_value("initialize", response.clone())?; + let mcp = parsed .capabilities .session .as_ref() - .and_then(|session| session.mcp.as_ref()); + .and_then(|s| s.mcp.as_ref()); DownstreamMcpMode::from_capabilities( - mcp.is_some_and(|mcp| mcp.http.is_some()), - mcp.is_some_and(|mcp| mcp.acp.is_some()), + mcp.is_some_and(|m| m.http.is_some()), + mcp.is_some_and(|m| m.acp.is_some()), ) } }; - if mode == DownstreamMcpMode::HttpAdapter { - self.advertise_native_mcp(response)?; - } - Ok(mode) - } - - fn advertise_native_mcp(self, response: &mut Value) -> Result<(), Error> { - let response = response - .as_object_mut() - .ok_or_else(|| invalid_initialize_response("result must be an object"))?; - match self { - Self::V1 => { - let mcp = response - .get_mut("agentCapabilities") - .and_then(Value::as_object_mut) - .and_then(|capabilities| capabilities.get_mut("mcpCapabilities")) - .and_then(Value::as_object_mut) - .ok_or_else(|| { - invalid_initialize_response( - "HTTP MCP support did not have an object capability container", - ) - })?; - mcp.insert("acp".into(), Value::Bool(true)); - } - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => { - let mcp = response - .get_mut("capabilities") - .and_then(Value::as_object_mut) - .and_then(|capabilities| capabilities.get_mut("session")) - .and_then(Value::as_object_mut) - .and_then(|session| session.get_mut("mcp")) - .and_then(Value::as_object_mut) - .ok_or_else(|| { - invalid_initialize_response( - "HTTP MCP support did not have an object capability container", - ) - })?; - mcp.insert("acp".into(), Value::Object(Map::new())); + let root = response.as_object_mut().ok_or_else(Error::invalid_params)?; + match self { + Self::V1 => { + let mcp = root + .get_mut("agentCapabilities") + .and_then(Value::as_object_mut) + .and_then(|capabilities| capabilities.get_mut("mcpCapabilities")) + .and_then(Value::as_object_mut) + .ok_or_else(Error::invalid_params)?; + mcp.insert("acp".into(), Value::Bool(true)); + } + #[cfg(feature = "unstable_protocol_v2")] + Self::V2 => { + let mcp = root + .get_mut("capabilities") + .and_then(Value::as_object_mut) + .and_then(|capabilities| capabilities.get_mut("session")) + .and_then(Value::as_object_mut) + .and_then(|session| session.get_mut("mcp")) + .and_then(Value::as_object_mut) + .ok_or_else(Error::invalid_params)?; + mcp.insert("acp".into(), Value::Object(Map::new())); + } } } - Ok(()) + Ok(mode) } pub(crate) fn is_session_setup_method(self, method: &str) -> bool { @@ -184,7 +185,7 @@ impl PolyfillProtocol { "session/fork" => { ForkSessionRequest::parse_message(request.method(), request.params())?; } - method => return Err(unexpected_session_setup_method(method)), + _ => return Err(Error::invalid_request().data("not a session setup method")), }, #[cfg(feature = "unstable_protocol_v2")] Self::V2 => match request.method() { @@ -198,7 +199,7 @@ impl PolyfillProtocol { "session/fork" => { v2::ForkSessionRequest::parse_message(request.method(), request.params())?; } - method => return Err(unexpected_session_setup_method(method)), + _ => return Err(Error::invalid_request().data("not a session setup method")), }, } Ok(()) @@ -206,165 +207,102 @@ impl PolyfillProtocol { pub(crate) fn native_server(self, value: Value) -> Option { let raw = value.as_object()?.clone(); - match self { + let (name, server_id) = match self { Self::V1 => { let McpServer::Acp(server) = serde_json::from_value(value).ok()? else { return None; }; - Some(NativeServer { - raw, - name: server.name, - server_id: server.server_id.to_string(), - }) + (server.name, server.server_id.to_string()) } #[cfg(feature = "unstable_protocol_v2")] Self::V2 => { let v2::McpServer::Acp(server) = serde_json::from_value(value).ok()? else { return None; }; - Some(NativeServer { - raw, - name: server.name, - server_id: server.server_id.to_string(), - }) + (server.name, server.server_id.to_string()) } - } - } - - pub(crate) fn connect_request(self, server_id: String) -> Result { - match self { - Self::V1 => ConnectMcpRequest::new(server_id).to_untyped_message(), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::ConnectMcpRequest::new(server_id).to_untyped_message(), - } - } - - pub(crate) fn connect_response_id(self, response: Value) -> Result { - match self { - Self::V1 => ConnectMcpResponse::from_value("mcp/connect", response) - .map(|response| response.connection_id.to_string()), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::ConnectMcpResponse::from_value("mcp/connect", response) - .map(|response| response.connection_id.to_string()), - } + }; + Some(NativeServer { + raw, + name, + server_id, + }) } pub(crate) fn message_request( self, - connection_id: String, - message: UntypedMessage, - ) -> Result { - let (method, params) = message.into_parts(); - let params = into_mcp_params(params)?; - match self { - Self::V1 => MessageMcpRequest::new(connection_id, method) - .params(params) - .to_untyped_message(), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::MessageMcpRequest::new(connection_id, method) - .params(params) - .to_untyped_message(), - } - } - - pub(crate) fn message_notification( - self, - connection_id: String, - message: UntypedMessage, + server_id: String, + request_id: String, + method: String, + params: Option>, + meta: Option, ) -> Result { - let (method, params) = message.into_parts(); - let params = into_mcp_params(params)?; - match self { - Self::V1 => MessageMcpNotification::new(connection_id, method) - .params(params) - .to_untyped_message(), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::MessageMcpNotification::new(connection_id, method) - .params(params) - .to_untyped_message(), + let mut wrapper = Map::new(); + wrapper.insert("serverId".into(), server_id.into()); + wrapper.insert("requestId".into(), request_id.into()); + wrapper.insert("method".into(), method.into()); + if let Some(params) = params { + wrapper.insert("params".into(), Value::Object(params)); } - } - - pub(crate) fn parse_message_request( - self, - request: UntypedMessage, - ) -> Result { - match self { - Self::V1 => { - let parsed = MessageMcpRequest::parse_message(request.method(), request.params())?; - Ok(NativeMcpMessage { - raw: request, - connection_id: parsed.connection_id.to_string(), - method: parsed.method, - params: parsed.params, - }) - } - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => { - let parsed = - v2::MessageMcpRequest::parse_message(request.method(), request.params())?; - Ok(NativeMcpMessage { - raw: request, - connection_id: parsed.connection_id.to_string(), - method: parsed.method, - params: parsed.params, - }) - } + if let Some(meta) = meta { + wrapper.insert("_meta".into(), meta); } - } - - pub(crate) fn parse_message_notification( - self, - notification: UntypedMessage, - ) -> Result { + let request = UntypedMessage { + method: "mcp/message".into(), + params: Value::Object(wrapper), + }; + // Validate the selected schema without losing unknown wrapper fields. match self { Self::V1 => { - let parsed = MessageMcpNotification::parse_message( - notification.method(), - notification.params(), + agent_client_protocol::schema::v1::MessageMcpRequest::parse_message( + request.method(), + request.params(), )?; - Ok(NativeMcpMessage { - raw: notification, - connection_id: parsed.connection_id.to_string(), - method: parsed.method, - params: parsed.params, - }) } #[cfg(feature = "unstable_protocol_v2")] Self::V2 => { - let parsed = v2::MessageMcpNotification::parse_message( - notification.method(), - notification.params(), - )?; - Ok(NativeMcpMessage { - raw: notification, - connection_id: parsed.connection_id.to_string(), - method: parsed.method, - params: parsed.params, - }) + v2::MessageMcpRequest::parse_message(request.method(), request.params())?; } } + Ok(request) } - pub(crate) fn disconnect_request(self, connection_id: String) -> Result { - match self { - Self::V1 => DisconnectMcpRequest::new(connection_id).to_untyped_message(), - #[cfg(feature = "unstable_protocol_v2")] - Self::V2 => v2::DisconnectMcpRequest::new(connection_id).to_untyped_message(), - } - } - - pub(crate) fn validate_disconnect_response(self, response: Value) -> Result<(), Error> { - match self { + pub(crate) fn parse_notification( + self, + raw: UntypedMessage, + ) -> Result { + let (server_id, request_id, method, params) = match self { Self::V1 => { - DisconnectMcpResponse::from_value("mcp/disconnect", response)?; + let parsed = + agent_client_protocol::schema::v1::MessageMcpNotification::parse_message( + raw.method(), + raw.params(), + )?; + ( + parsed.server_id.to_string(), + parsed.request_id.to_string(), + parsed.method, + parsed.params, + ) } #[cfg(feature = "unstable_protocol_v2")] Self::V2 => { - v2::DisconnectMcpResponse::from_value("mcp/disconnect", response)?; + let parsed = v2::MessageMcpNotification::parse_message(raw.method(), raw.params())?; + ( + parsed.server_id.to_string(), + parsed.request_id.to_string(), + parsed.method, + parsed.params, + ) } - } - Ok(()) + }; + Ok(NativeMcpNotification { + raw, + server_id, + request_id, + method, + params, + }) } } @@ -401,16 +339,19 @@ impl NativeServer { mut self, protocol: PolyfillProtocol, url: String, + token: &str, ) -> Result { self.raw.remove("serverId"); - self.raw.insert("type".into(), Value::String("http".into())); - self.raw.insert("name".into(), Value::String(self.name)); - self.raw.insert("url".into(), Value::String(url)); - // V1 requires the field and v2 accepts it. Keeping the explicit empty - // list gives both versions one stable raw compatibility shape. - self.raw.insert("headers".into(), Value::Array(Vec::new())); + self.raw.insert("type".into(), "http".into()); + self.raw.insert("name".into(), self.name.into()); + self.raw.insert("url".into(), url.into()); + self.raw.insert( + "headers".into(), + serde_json::json!([ + { "name": "Authorization", "value": format!("Bearer {token}") } + ]), + ); let declaration = Value::Object(self.raw); - match protocol { PolyfillProtocol::V1 => { serde_json::from_value::(declaration.clone()) @@ -426,272 +367,10 @@ impl NativeServer { } } -#[derive(Debug)] -pub(crate) struct NativeMcpMessage { +pub(crate) struct NativeMcpNotification { pub(crate) raw: UntypedMessage, - pub(crate) connection_id: String, + pub(crate) server_id: String, + pub(crate) request_id: String, pub(crate) method: String, pub(crate) params: Option>, } - -pub(crate) fn native_params_into_value(params: Option>) -> Value { - params.map_or(Value::Null, Value::Object) -} - -fn into_mcp_params(params: Value) -> Result>, Error> { - match params { - Value::Null => Ok(None), - Value::Object(params) => Ok(Some(params)), - params => Err(Error::invalid_params().data(serde_json::json!({ - "reason": "MCP message params must be an object or null", - "params": params, - }))), - } -} - -fn invalid_initialize_protocol_version() -> Error { - Error::invalid_params().data("initialize.protocolVersion must be a valid ACP protocol version") -} - -fn unsupported_protocol_version(version: ProtocolVersion) -> Error { - Error::invalid_request().data(format!( - "MCP-over-ACP polyfill does not support ACP protocol version {version}" - )) -} - -fn unexpected_session_setup_method(method: &str) -> Error { - Error::invalid_request().data(format!( - "`{method}` is not a session setup method for the selected ACP version" - )) -} - -fn invalid_initialize_response(reason: &'static str) -> Error { - Error::invalid_params().data(format!("invalid initialize response: {reason}")) -} - -#[cfg(test)] -mod tests { - use agent_client_protocol::{ - JsonRpcMessage, - schema::{ProtocolVersion, v1}, - }; - - #[cfg(feature = "unstable_protocol_v2")] - use agent_client_protocol::{ErrorCode, JsonRpcResponse}; - - use super::PolyfillProtocol; - - #[test] - fn http_declaration_preserves_extension_fields() { - let declaration = serde_json::json!({ - "type": "acp", - "name": "native", - "serverId": "native-id", - "_meta": { - "source": "test" - }, - "futureField": { - "preserve": true - } - }); - let native = PolyfillProtocol::V1 - .native_server(declaration) - .expect("the declaration should be recognized as native MCP"); - - let transformed = native - .http_declaration(PolyfillProtocol::V1, "http://127.0.0.1:4321".to_string()) - .expect("the transformed declaration should be valid v1 MCP"); - - assert_eq!( - transformed, - serde_json::json!({ - "type": "http", - "name": "native", - "url": "http://127.0.0.1:4321", - "headers": [], - "_meta": { - "source": "test" - }, - "futureField": { - "preserve": true - } - }) - ); - } - - #[test] - fn native_message_keeps_the_original_wrapper() { - let request = agent_client_protocol::UntypedMessage { - method: "mcp/message".to_string(), - params: serde_json::json!({ - "connectionId": "connection", - "method": "tools/list", - "params": { - "cursor": "next" - }, - "_meta": { - "trace": "preserve" - }, - "futureField": true - }), - }; - - let parsed = PolyfillProtocol::V1 - .parse_message_request(request.clone()) - .expect("the native wrapper should parse"); - - assert_eq!(parsed.raw, request); - assert_eq!(parsed.connection_id, "connection"); - assert_eq!(parsed.method, "tools/list"); - assert_eq!( - parsed.params, - Some(serde_json::Map::from_iter([( - "cursor".to_string(), - serde_json::json!("next") - )])) - ); - } - - #[test] - fn v1_session_setup_methods_match_the_stable_schema() { - assert!(PolyfillProtocol::V1.is_session_setup_method("session/new")); - assert!(PolyfillProtocol::V1.is_session_setup_method("session/load")); - assert!(PolyfillProtocol::V1.is_session_setup_method("session/resume")); - assert_eq!( - PolyfillProtocol::V1.is_session_setup_method("session/fork"), - cfg!(feature = "unstable_session_fork") - ); - assert!(!PolyfillProtocol::V1.is_session_setup_method("session/prompt")); - } - - #[test] - fn session_setup_validation_allows_extensions_but_rejects_invalid_fields() { - let mut request = v1::NewSessionRequest::new(std::path::PathBuf::from("/tmp")) - .to_untyped_message() - .expect("the session request should serialize"); - request.params["futureField"] = serde_json::json!({ - "preserve": true - }); - PolyfillProtocol::V1 - .validate_session_setup_request(&request) - .expect("extension fields should remain forward-compatible"); - - request.params["cwd"] = serde_json::json!(42); - let error = PolyfillProtocol::V1 - .validate_session_setup_request(&request) - .expect_err("invalid selected-schema fields must be rejected"); - assert_eq!(error.code, agent_client_protocol::ErrorCode::InvalidParams); - } - - #[cfg(feature = "unstable_protocol_v2")] - #[test] - fn v2_session_setup_methods_exclude_v1_load() { - assert!(PolyfillProtocol::V2.is_session_setup_method("session/new")); - assert!(!PolyfillProtocol::V2.is_session_setup_method("session/load")); - assert!(PolyfillProtocol::V2.is_session_setup_method("session/resume")); - assert_eq!( - PolyfillProtocol::V2.is_session_setup_method("session/fork"), - cfg!(feature = "unstable_session_fork") - ); - assert!(!PolyfillProtocol::V2.is_session_setup_method("session/prompt")); - } - - #[cfg(feature = "unstable_protocol_v2")] - #[test] - fn future_protocol_version_is_not_assumed_to_be_v2() { - let initialize = agent_client_protocol::schema::v2::InitializeRequest::new( - ProtocolVersion::V2, - agent_client_protocol::schema::v2::Implementation::new("test", "1.0.0"), - ); - let mut request = - agent_client_protocol::schema::v2::InitializeProxyRequest::new(initialize) - .to_untyped_message() - .expect("the initialize request should serialize"); - request.params["protocolVersion"] = serde_json::json!(3); - - let error = PolyfillProtocol::from_initialize_request(&request) - .expect_err("an unselected future schema must not be interpreted as v2"); - - assert_eq!(error.code, ErrorCode::InvalidRequest); - assert_eq!( - error.data, - Some(serde_json::json!( - "MCP-over-ACP polyfill does not support ACP protocol version 3" - )) - ); - } - - #[cfg(feature = "unstable_protocol_v2")] - #[test] - fn v2_initialize_adaptation_preserves_the_raw_response() { - use agent_client_protocol::schema::v2; - - let response = v2::InitializeResponse::new( - ProtocolVersion::V2, - v2::Implementation::new("test", "1.0.0"), - ) - .capabilities( - v2::AgentCapabilities::new().session( - v2::SessionCapabilities::new() - .mcp(v2::McpCapabilities::new().http(v2::McpHttpCapabilities::new())), - ), - ); - let mut response = serde_json::to_value(response).expect("the response should serialize"); - response["futureField"] = serde_json::json!({ - "preserve": true - }); - - let mode = PolyfillProtocol::V2 - .transform_initialize_response(&mut response) - .expect("the v2 HTTP capability should be adaptable"); - - assert_eq!(mode, super::DownstreamMcpMode::HttpAdapter); - assert_eq!( - response["capabilities"]["session"]["mcp"]["acp"], - serde_json::json!({}) - ); - assert_eq!( - response["futureField"], - serde_json::json!({ - "preserve": true - }) - ); - - v2::InitializeResponse::from_value("initialize", response) - .expect("the adapted response should remain valid v2"); - } - - #[cfg(feature = "unstable_protocol_v2")] - #[test] - fn v2_disconnect_uses_and_validates_the_selected_schema() { - use agent_client_protocol::schema::v2; - - let request = PolyfillProtocol::V2 - .disconnect_request("connection".to_string()) - .expect("the v2 disconnect request should serialize"); - let parsed = v2::DisconnectMcpRequest::parse_message(request.method(), request.params()) - .expect("the disconnect request should be valid v2"); - assert_eq!(parsed.connection_id.to_string(), "connection"); - - let response = serde_json::to_value(v2::DisconnectMcpResponse::new()) - .expect("the v2 disconnect response should serialize"); - PolyfillProtocol::V2 - .validate_disconnect_response(response) - .expect("the v2 disconnect response should validate"); - } - - #[test] - fn v1_initialize_request_selects_v1() { - let request = agent_client_protocol::schema::InitializeProxyRequest { - initialize: v1::InitializeRequest::new(ProtocolVersion::V1), - } - .to_untyped_message() - .expect("the initialize request should serialize"); - - assert_eq!( - PolyfillProtocol::from_initialize_request(&request) - .expect("the request should select v1"), - PolyfillProtocol::V1 - ); - } -} diff --git a/src/agent-client-protocol-rmcp/Cargo.toml b/src/agent-client-protocol-rmcp/Cargo.toml index 4952b98d..ab207164 100644 --- a/src/agent-client-protocol-rmcp/Cargo.toml +++ b/src/agent-client-protocol-rmcp/Cargo.toml @@ -13,11 +13,16 @@ categories = ["development-tools"] [features] default = [] unstable_mcp_over_acp = ["agent-client-protocol/unstable_mcp_over_acp"] +unstable_protocol_v2 = ["agent-client-protocol/unstable_protocol_v2"] [[example]] name = "with_mcp_server" required-features = ["unstable_mcp_over_acp"] +[[example]] +name = "stateless_native_mcp" +required-features = ["unstable_mcp_over_acp", "unstable_protocol_v2"] + [dependencies] agent-client-protocol = { workspace = true, features = ["schemars"] } futures.workspace = true diff --git a/src/agent-client-protocol-rmcp/README.md b/src/agent-client-protocol-rmcp/README.md index 4b12c962..7cee740c 100644 --- a/src/agent-client-protocol-rmcp/README.md +++ b/src/agent-client-protocol-rmcp/README.md @@ -9,12 +9,20 @@ runtime-agnostic MCP server framework from `agent-client-protocol`. It lets you Rust, serve them directly, or attach them to an ACP proxy. Attached servers are advertised with the opt-in native MCP-over-ACP transport: -`McpServer::Acp` plus `mcp/connect`, `mcp/message`, and `mcp/disconnect`. This +`McpServer::Acp` plus request-scoped `mcp/message` operations targeting MCP +2026-07-28. There is no MCP initialization or connect/disconnect exchange. This crate does not enable the core SDK's `unstable_mcp_over_acp` feature merely to build or directly serve a server. Enable this crate's matching `unstable_mcp_over_acp` feature when using `with_mcp_server`. Use -`agent-client-protocol-polyfill` when the final agent accepts HTTP but not -ACP-transport MCP servers. +`agent-client-protocol-polyfill` when the final agent has a modern MCP HTTP +client but does not consume ACP-transport MCP servers natively. + +For a direct ACP client/agent example using real rmcp tools, run: + +```sh +cargo run -p agent-client-protocol-rmcp --example stateless_native_mcp \ + --features unstable_mcp_over_acp,unstable_protocol_v2 +``` ## Usage diff --git a/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs b/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs new file mode 100644 index 00000000..8263459f --- /dev/null +++ b/src/agent-client-protocol-rmcp/examples/stateless_native_mcp.rs @@ -0,0 +1,129 @@ +//! Run with `cargo run -p agent-client-protocol-rmcp --example stateless_native_mcp +//! --features unstable_mcp_over_acp,unstable_protocol_v2`. +//! No MCP initialize or separate MCP transport: the client attaches an rmcp service +//! to an ACP session and the agent invokes it through `mcp/message`. + +use agent_client_protocol::{ + Agent, Client, Error, Responder, V2ConnectionTo, + mcp_server::McpServer, + schema::{ProtocolVersion, v2}, +}; +use agent_client_protocol_rmcp::McpServerExt; +use rmcp::{ + ErrorData, RoleServer, ServerHandler, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, ServerCapabilities, ServerConfig, + }, + service::RequestContext, +}; +use serde_json::json; +use std::sync::{Arc, Mutex}; +use tokio::sync::oneshot; + +struct Echo; + +impl ServerHandler for Echo { + fn get_info(&self) -> ServerConfig { + ServerConfig::new(ServerCapabilities::builder().enable_tools().build()) + } + + fn call_tool( + &self, + params: CallToolRequestParams, + _cx: RequestContext, + ) -> impl std::future::Future> + Send { + std::future::ready(if params.name == "echo" { + Ok(CallToolResult::structured(json!({"echoed": params.arguments})).into()) + } else { + Err(ErrorData::invalid_params("unknown tool", None)) + }) + } +} + +#[tokio::main] +async fn main() -> Result<(), Error> { + let (done_tx, done_rx) = oneshot::channel(); + let done_tx = Arc::new(Mutex::new(Some(done_tx))); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("echo-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new() + .session(v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + )), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let server_id = match request.mcp_servers.as_slice() { + [v2::McpServer::Acp(server)] if server.name == "echo" => { + server.server_id.clone() + } + other => panic!("unexpected MCP declaration: {other:?}"), + }; + let done_tx = done_tx.lock().unwrap().take().expect("one session"); + let call_cx = cx.clone(); + cx.spawn(async move { + let mut params = json!({"name": "echo", "arguments": {"message": "hello ACP"}}); + params["_meta"] = json!({ + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + "io.modelcontextprotocol/clientInfo": {"name": "echo-agent", "version": "1"} + }); + let response = call_cx + .send_request( + v2::MessageMcpRequest::new(server_id, "echo-1", "tools/call") + .params(params.as_object().expect("object params").clone()), + ) + .block_task() + .await; + drop(done_tx.send(response)); + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new( + "echo-session", + ))) + }, + agent_client_protocol::on_receive_request!(), + ); + + Client + .v2() + .connect_with(agent, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("echo-client", "1"), + )) + .block_task() + .await?; + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(McpServer::::from_rmcp("echo", || Echo))? + .start_session() + .block_task() + .await?; + let response = done_rx.await.map_err(Error::into_internal_error)??; + match response { + v2::MessageMcpResponse::Result { result, .. } => println!("{result}"), + v2::MessageMcpResponse::Error { error, .. } => { + eprintln!("MCP error {}: {}", error.code, error.message); + } + _ => return Err(Error::internal_error().data("unknown MCP carrier")), + } + Ok(()) + }) + .await +} diff --git a/src/agent-client-protocol-rmcp/src/builder.rs b/src/agent-client-protocol-rmcp/src/builder.rs index 585cd94c..6f8d2af4 100644 --- a/src/agent-client-protocol-rmcp/src/builder.rs +++ b/src/agent-client-protocol-rmcp/src/builder.rs @@ -15,6 +15,8 @@ use schemars::JsonSchema; use serde::{Serialize, de::DeserializeOwned}; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; +#[cfg(feature = "unstable_mcp_over_acp")] +use acp::mcp_server::{McpOutcome, McpRequest, McpRequestContext, McpService}; use agent_client_protocol as acp; use agent_client_protocol::{ ByteStreams, ChainRun, ConnectTo, DynConnectTo, NullRun, RunWithConnectionTo, @@ -231,13 +233,22 @@ where /// feature, it can also be attached through /// `SessionBuilder::with_mcp_server` or `Builder::with_mcp_server`. pub fn build(self) -> McpServer { - McpServer::new( - McpServerBuilt { - name: self.name, - data: Arc::new(self.data), - }, - self.runner, - ) + let built = McpServerBuilt { + name: self.name, + data: Arc::new(self.data), + }; + #[cfg(feature = "unstable_mcp_over_acp")] + { + let standalone = McpServerBuilt { + name: built.name.clone(), + data: built.data.clone(), + }; + McpServer::new_service_with_standalone(built, standalone, self.runner) + } + #[cfg(not(feature = "unstable_mcp_over_acp"))] + { + McpServer::new(built, self.runner) + } } } @@ -246,6 +257,21 @@ struct McpServerBuilt { data: Arc>, } +#[cfg(feature = "unstable_mcp_over_acp")] +impl McpService for McpServerBuilt { + fn execute( + &self, + request: McpRequest, + context: McpRequestContext, + ) -> BoxFuture<'static, Result> { + let handler = McpServerConnection { + data: self.data.clone(), + mcp_connection: context.connection().clone(), + }; + crate::native::execute(Arc::new(handler), request, context) + } +} + impl McpServerConnect for McpServerBuilt { fn name(&self) -> String { self.name.clone() diff --git a/src/agent-client-protocol-rmcp/src/lib.rs b/src/agent-client-protocol-rmcp/src/lib.rs index 1a91a54d..34f1f5e2 100644 --- a/src/agent-client-protocol-rmcp/src/lib.rs +++ b/src/agent-client-protocol-rmcp/src/lib.rs @@ -40,13 +40,21 @@ //! ``` use agent_client_protocol::mcp_server::{McpConnectionTo, McpServer, McpServerConnect}; +#[cfg(feature = "unstable_mcp_over_acp")] +use agent_client_protocol::mcp_server::{McpOutcome, McpRequest, McpRequestContext, McpService}; use agent_client_protocol::role; use agent_client_protocol::{ByteStreams, ConnectTo, DynConnectTo, NullRun, Role}; +#[cfg(feature = "unstable_mcp_over_acp")] +use futures::future::BoxFuture; use futures_concurrency::future::TryJoin as _; use rmcp::ServiceExt; +#[cfg(feature = "unstable_mcp_over_acp")] +use std::sync::Arc; use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; mod builder; +#[cfg(feature = "unstable_mcp_over_acp")] +mod native; pub use agent_client_protocol::mcp_server::{EnabledTools, McpTool}; pub use agent_client_protocol::{tool_fn, tool_fn_mut}; @@ -76,6 +84,29 @@ pub trait McpServerExt { new_fn: F, } + #[cfg(feature = "unstable_mcp_over_acp")] + struct SharedRmcp { + new_fn: Arc, + service: std::sync::OnceLock>, + } + + #[cfg(feature = "unstable_mcp_over_acp")] + impl McpService for SharedRmcp + where + Counterpart: Role, + F: Fn() -> S + Send + Sync + 'static, + S: rmcp::Service, + { + fn execute( + &self, + request: McpRequest, + context: McpRequestContext, + ) -> BoxFuture<'static, Result> { + let service = self.service.get_or_init(|| Arc::new((self.new_fn)())); + native::execute(service.clone(), request, context) + } + } + impl McpServerConnect for RmcpServer where Counterpart: Role, @@ -95,13 +126,34 @@ pub trait McpServerExt { } } - McpServer::new( - RmcpServer { - name: name.to_string(), - new_fn, - }, - NullRun, - ) + #[cfg(feature = "unstable_mcp_over_acp")] + { + // Feature unification must not construct an unused ACP service + // when this server is only used through its standalone adapter. + let new_fn = Arc::new(new_fn); + let shared = SharedRmcp { + new_fn: new_fn.clone(), + service: std::sync::OnceLock::new(), + }; + McpServer::new_service_with_standalone( + shared, + RmcpServer { + name: name.to_string(), + new_fn: move || new_fn(), + }, + NullRun, + ) + } + #[cfg(not(feature = "unstable_mcp_over_acp"))] + { + McpServer::new( + RmcpServer { + name: name.to_string(), + new_fn, + }, + NullRun, + ) + } } } @@ -130,10 +182,7 @@ where let byte_streams = ByteStreams::new(mcp_client_write.compat_write(), mcp_client_read.compat()); - // Spawn task to connect byte_streams to the provided client - drop(ConnectTo::::connect_to(byte_streams, client).await); - - Ok(()) + ConnectTo::::connect_to(byte_streams, client).await }; let bytes_to_rmcp = async { diff --git a/src/agent-client-protocol-rmcp/src/native.rs b/src/agent-client-protocol-rmcp/src/native.rs new file mode 100644 index 00000000..6c5cc9ba --- /dev/null +++ b/src/agent-client-protocol-rmcp/src/native.rs @@ -0,0 +1,239 @@ +//! Direct, request-scoped rmcp transport for ACP (no byte-stream emulation). + +use std::{ + future::Future, + sync::{Arc, Mutex}, +}; + +use acp::{ + Role, + mcp_server::{MCP_BACKEND_FAILURE, McpOutcome, McpRequest, McpRequestContext}, +}; +use agent_client_protocol as acp; +use futures::{ + channel::oneshot, + future::{BoxFuture, Either}, +}; +use rmcp::{ + RoleServer, Service, + model::ClientJsonRpcMessage, + service::{self, NotificationContext, RequestContext}, + transport::OneshotTransport, +}; +use tokio_util::sync::CancellationToken; + +/// Reuses the same application service but owns every handler future and its +/// cancellation on this one operation. +struct OperationService { + app: Arc, + cancel: CancellationToken, + completions: Arc>>>, +} + +impl> Service for OperationService { + fn handle_request( + &self, + request: ::PeerReq, + context: RequestContext, + ) -> impl Future::Resp, rmcp::ErrorData>> + + Send + + '_ { + let (done_tx, done_rx) = oneshot::channel(); + self.completions + .lock() + .expect("MCP operation poisoned") + .push(done_rx); + let cancel = self.cancel.clone(); + async move { + let result = tokio::select! { + biased; + () = cancel.cancelled() => Err(rmcp::ErrorData::internal_error("operation cancelled", None)), + result = self.app.handle_request(request, context) => result, + }; + let _sent = done_tx.send(()); + result + } + } + + fn handle_notification( + &self, + notification: ::PeerNot, + context: NotificationContext, + ) -> impl Future> + Send + '_ { + let (done_tx, done_rx) = oneshot::channel(); + self.completions + .lock() + .expect("MCP operation poisoned") + .push(done_rx); + let cancel = self.cancel.clone(); + async move { + let result = tokio::select! { + biased; + () = cancel.cancelled() => Ok(()), + result = self.app.handle_notification(notification, context) => result, + }; + let _sent = done_tx.send(()); + result + } + } + + fn get_info(&self) -> ::Info { + self.app.get_info() + } + + fn supported_protocol_versions( + &self, + ) -> std::borrow::Cow<'static, [rmcp::model::ProtocolVersion]> { + self.app.supported_protocol_versions() + } +} + +/// Execute one request against shared rmcp application state. Neither the +/// client transport nor the rmcp server's actor is allowed to escape this call. +pub(crate) fn execute( + app: Arc, + request: McpRequest, + context: McpRequestContext, +) -> BoxFuture<'static, Result> +where + R: Role, + S: Service, +{ + Box::pin(async move { + let id = context.request_id().0.to_string(); + let raw = serde_json::json!({ + "jsonrpc": "2.0", + "id": id, + "method": request.method, + "params": request.params, + }); + let inbound: ClientJsonRpcMessage = match serde_json::from_value(raw) { + Ok(request) => request, + Err(error) => { + return Ok(McpOutcome::Error( + acp::schema::v1::McpError::new( + if error.to_string().contains("unknown variant") { + -32601 + } else { + -32602 + }, + "Invalid MCP request", + ) + .data(serde_json::Value::String(error.to_string())), + )); + } + }; + let (transport, mut output) = OneshotTransport::::new(inbound); + let cancel = CancellationToken::new(); + let completions = Arc::new(Mutex::new(Vec::new())); + let handler = OperationService { + app, + cancel: cancel.clone(), + completions: completions.clone(), + }; + let mut running = service::serve_directly_with_ct(handler, transport, None, cancel.clone()); + let operation = async { + while let Some(outbound) = output.recv().await { + let value = serde_json::to_value(outbound).map_err(|error| { + acp::Error::new( + acp::mcp_server::MCP_BACKEND_FAILURE, + format!("cannot serialize MCP output: {error}"), + ) + })?; + match value { + serde_json::Value::Object(mut object) if object.contains_key("method") => { + if object.contains_key("id") { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "reverse MCP requests are not supported", + )); + } + let method = object + .remove("method") + .and_then(|v| v.as_str().map(str::to_owned)) + .ok_or_else(|| { + acp::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend notification has no method", + ) + })?; + let params = match object.remove("params") { + None | Some(serde_json::Value::Null) => None, + Some(serde_json::Value::Object(params)) => Some(params), + _ => { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend notification parameters must be an object", + )); + } + }; + context.send_notification(method, params).await?; + } + serde_json::Value::Object(mut object) if object.contains_key("result") => { + if object.get("id").and_then(serde_json::Value::as_str) != Some(id.as_str()) + { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "MCP response ID mismatch", + )); + } + return Ok(McpOutcome::Result( + object.remove("result").expect("checked result"), + )); + } + serde_json::Value::Object(mut object) if object.contains_key("error") => { + if object.get("id").and_then(serde_json::Value::as_str) != Some(id.as_str()) + { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "MCP error ID mismatch", + )); + } + let error = + serde_json::from_value(object.remove("error").expect("checked error")) + .map_err(|error| { + acp::Error::new( + acp::mcp_server::MCP_BACKEND_FAILURE, + format!("invalid MCP error from backend: {error}"), + ) + })?; + return Ok(McpOutcome::Error(error)); + } + _ => { + return Err(acp::Error::new( + MCP_BACKEND_FAILURE, + "unexpected MCP output", + )); + } + } + } + Err(acp::Error::new( + acp::mcp_server::MCP_BACKEND_FAILURE, + "MCP backend closed without a response", + )) + }; + let cancelled = async { + let acp = context.cancellation().cancelled(); + let operation = context.operation_cancellation().cancelled(); + futures::pin_mut!(acp, operation); + let _reason = futures::future::select(acp, operation).await; + }; + let result = match futures::future::select(Box::pin(operation), Box::pin(cancelled)).await { + Either::Left((result, _)) => result, + Either::Right(((), _)) => Err(acp::Error::request_cancelled()), + }; + cancel.cancel(); + let closed = running.close().await; + let handlers = std::mem::take(&mut *completions.lock().expect("MCP operation poisoned")); + for completion in handlers { + let _finished = completion.await; + } + closed.map_err(|error| { + acp::Error::new( + acp::mcp_server::MCP_BACKEND_FAILURE, + format!("MCP backend cleanup failed: {error}"), + ) + })?; + result + }) +} diff --git a/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs b/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs new file mode 100644 index 00000000..5adfcdc3 --- /dev/null +++ b/src/agent-client-protocol-rmcp/tests/stateless_native_mcp.rs @@ -0,0 +1,516 @@ +//! Native ACP attachment of an rmcp service (not standalone MCP transport). +#![cfg(all(feature = "unstable_protocol_v2", feature = "unstable_mcp_over_acp"))] + +use std::{ + collections::HashMap, + future::Future, + sync::{Arc, Mutex}, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, Channel, Client, Error, Responder, V2ConnectionTo, + mcp_server::McpServer, + schema::{ProtocolVersion, v2}, +}; +use agent_client_protocol_rmcp::McpServerExt; +use rmcp::{ + ErrorData, RoleServer, ServerHandler, + model::{ + CallToolRequestParams, CallToolResponse, CallToolResult, InputRequiredResult, + ServerCapabilities, ServerConfig, SubscriptionFilter, + }, + service::{RequestContext, SubscriptionContext}, +}; +use serde_json::{Value, json}; +use tokio::sync::{mpsc, oneshot}; + +fn meta(marker: &str) -> Value { + json!({"io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {"elicitation": {"form": {}}}, + "io.modelcontextprotocol/clientInfo": {"name": "native-acp", "version": "1"}, + "example/marker": marker}) +} + +async fn message( + cx: &V2ConnectionTo, + server: &v2::McpServerAcpId, + id: &str, + method: &str, + mut params: Value, + marker: &str, +) -> Result { + params["_meta"] = meta(marker); + let response = cx + .send_request( + v2::MessageMcpRequest::new(server.clone(), id.to_owned(), method) + .params(params.as_object().expect("object params").clone()), + ) + .block_task() + .await?; + match response { + v2::MessageMcpResponse::Result { result, .. } => Ok(result), + v2::MessageMcpResponse::Error { error, .. } => Err(Error::new(error.code, error.message) + .data(match error.data { + agent_client_protocol::schema::MaybeUndefined::Value(value) => Some(value), + agent_client_protocol::schema::MaybeUndefined::Null => Some(Value::Null), + agent_client_protocol::schema::MaybeUndefined::Undefined => None, + })), + _ => Err(Error::internal_error().data("unexpected MCP carrier outcome")), + } +} + +struct DropSignal(Arc>>>); +impl Drop for DropSignal { + fn drop(&mut self) { + if let Some(tx) = self.0.lock().unwrap().take() { + let _ = tx.send(()); + } + } +} + +struct Service { + _drop: DropSignal, + started: Arc>>>, + stopped: Arc>>>, + pending: Arc, oneshot::Sender<()>)>>>, +} +impl ServerHandler for Service { + fn get_info(&self) -> ServerConfig { + ServerConfig::new( + ServerCapabilities::builder() + .enable_tools() + .enable_tool_list_changed() + .build(), + ) + } + fn call_tool( + &self, + request: CallToolRequestParams, + cx: RequestContext, + ) -> impl Future> + Send { + let pending = if request.name.as_ref() == "hang" { + let probe = request + .arguments + .as_ref() + .and_then(|args| args.get("probe")) + .and_then(Value::as_str) + .expect("pending tool requires a named probe"); + Some( + self.pending + .lock() + .unwrap() + .remove(probe) + .expect("distinct operation probe"), + ) + } else { + None + }; + async move { + if let Some((started, dropped)) = pending { + let _drop = DropSignal(Arc::new(Mutex::new(Some(dropped)))); + let _started = started.send(()); + // Deliberately ignore rmcp RequestContext::ct: the adapter must + // drop this future on outer cancellation and join its cleanup. + std::future::pending::<()>().await; + } + match request.name.as_ref() { + "retry" if request.request_state.is_none() => { + let inputs = serde_json::from_value(json!({"confirmation": { + "method": "elicitation/create", "params": {"mode": "form", + "message": "Confirm", "requestedSchema": {"type": "object", + "properties": {"approved": {"type": "boolean"}}}} + }})) + .expect("valid elicitation"); + Ok(InputRequiredResult::new(Some(inputs), Some("retry-state".into())).into()) + } + "retry" if request.request_state.as_deref() == Some("retry-state") => Ok( + CallToolResult::structured(json!({"marker": cx.meta.get("example/marker"), + "responses": request.input_responses})) + .into(), + ), + "echo" => Ok(CallToolResult::structured( + json!({"marker": cx.meta.get("example/marker")}), + ) + .into()), + _ => Err(ErrorData::invalid_params( + "unknown tool or state", + Some(json!({"source": "rmcp"})), + )), + } + } + } + fn accepted_subscription_filter( + &self, + requested: &SubscriptionFilter, + ) -> Option { + Some(requested.clone()) + } + async fn listen(&self, cx: SubscriptionContext) -> Result<(), ErrorData> { + let _stopped = DropSignal(self.stopped.clone()); + cx.sink() + .notify_tool_list_changed() + .await + .map_err(|e| ErrorData::internal_error(e.to_string(), None))?; + if let Some(tx) = self.started.lock().unwrap().take() { + let _ = tx.send(()); + } + cx.cancelled().await; + Ok(()) + } +} + +async fn exercise( + cx: V2ConnectionTo, + server: v2::McpServerAcpId, + started: oneshot::Receiver<()>, + stopped: oneshot::Receiver<()>, + pending: Vec<(String, oneshot::Receiver<()>, oneshot::Receiver<()>)>, +) -> Result { + let direct = message( + &cx, + &server, + "direct-1", + "tools/call", + json!({"name": "echo", "arguments": {}}), + "direct", + ) + .await?; + assert_eq!(direct["structuredContent"]["marker"], "direct"); + let discovered = message( + &cx, + &server, + "discover-1", + "server/discover", + json!({}), + "discover", + ) + .await?; + assert!( + discovered["supportedVersions"] + .as_array() + .unwrap() + .contains(&json!("2026-07-28")) + ); + let first = message( + &cx, + &server, + "retry-1", + "tools/call", + json!({"name": "retry", "arguments": {}}), + "first", + ) + .await?; + assert_eq!(first["resultType"], "input_required"); + assert_eq!( + first["inputRequests"]["confirmation"]["method"], + "elicitation/create" + ); + let responses = json!({"confirmation": {"action": "accept", "content": {"approved": true}}}); + let retry = message( + &cx, + &server, + "retry-2", + "tools/call", + json!({"name": "retry", "arguments": {}, "requestState": first["requestState"], + "inputResponses": responses}), + "second", + ) + .await?; + assert_eq!(retry["structuredContent"]["marker"], "second"); + assert_eq!(retry["structuredContent"]["responses"], responses); + let error = message( + &cx, + &server, + "error-1", + "tools/call", + json!({"name": "missing", "arguments": {}}), + "error", + ) + .await + .expect_err("rmcp error"); + assert_eq!( + serde_json::to_value(error)?["data"], + json!({"source": "rmcp"}) + ); + let mut params = json!({"notifications": {"toolsListChanged": true}}); + params["_meta"] = meta("listen"); + let subscription = cx.send_request( + v2::MessageMcpRequest::new(server.clone(), "listen-1", "subscriptions/listen") + .params(params.as_object().expect("object params").clone()), + ); + started.await.map_err(Error::into_internal_error)?; + let parallel = message( + &cx, + &server, + "parallel-1", + "tools/call", + json!({"name": "echo", "arguments": {}}), + "parallel", + ) + .await?; + assert_eq!(parallel["structuredContent"]["marker"], "parallel"); + subscription.cancel()?; + stopped.await.map_err(Error::into_internal_error)?; + for (index, (probe, started, dropped)) in pending.into_iter().enumerate() { + let id = format!("hang-{index}"); + let mut params = json!({"name": "hang", "arguments": {"probe": probe}}); + params["_meta"] = meta(&id); + let request = cx.send_request( + v2::MessageMcpRequest::new(server.clone(), id.clone(), "tools/call") + .params(params.as_object().expect("object params").clone()), + ); + started.await.map_err(Error::into_internal_error)?; + request.cancel()?; + // This is a distinct operation-local future, not the shared service's + // destructor. Cleanup must precede the cancellation response. + dropped.await.map_err(Error::into_internal_error)?; + let error = request + .block_task() + .await + .expect_err("cancelled MCP request"); + assert_eq!(i32::from(error.code), -32800); + let healthy = message( + &cx, + &server, + &format!("healthy-{index}"), + "tools/call", + json!({"name": "echo", "arguments": {}}), + &id, + ) + .await?; + assert_eq!(healthy["structuredContent"]["marker"], id); + } + Ok(server) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn native_acp_stateless_rmcp_lifecycle() -> Result<(), Error> { + tokio::time::timeout(Duration::from_secs(10), async { + let (start_tx, start_rx) = oneshot::channel(); + let (stop_tx, stop_rx) = oneshot::channel(); + let (drop_tx, drop_rx) = oneshot::channel(); + let (result_tx, result_rx) = oneshot::channel(); + let mut pending_checks = Vec::new(); + let mut pending_handlers = HashMap::new(); + for name in ["first", "second"] { + let (started_tx, started_rx) = oneshot::channel(); + let (dropped_tx, dropped_rx) = oneshot::channel(); + pending_handlers.insert(name.to_owned(), (started_tx, dropped_tx)); + pending_checks.push((name.to_owned(), started_rx, dropped_rx)); + } + let pending_handlers = Arc::new(Mutex::new(pending_handlers)); + let invocation = Arc::new(Mutex::new(Some(( + start_rx, + stop_rx, + pending_checks, + result_tx, + )))); + let (notifications_tx, mut notifications_rx) = mpsc::unbounded_channel(); + let started = Arc::new(Mutex::new(Some(start_tx))); + let stopped = Arc::new(Mutex::new(Some(stop_tx))); + let dropped = Arc::new(Mutex::new(Some(drop_tx))); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("native-rmcp-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + ), + ), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let server = match request.mcp_servers.as_slice() { + [v2::McpServer::Acp(server)] if server.name == "real-rmcp" => { + server.server_id.clone() + } + other => panic!("unexpected declarations: {other:?}"), + }; + let (start_rx, stop_rx, pending_checks, result_tx) = + invocation.lock().unwrap().take().expect("one session"); + let call_cx = cx.clone(); + cx.spawn(async move { + let result = + exercise(call_cx, server, start_rx, stop_rx, pending_checks).await; + drop(result_tx.send(result)); + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new( + "native-session", + ))) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_notification( + async move |notification: v2::MessageMcpNotification, + _cx: V2ConnectionTo| { + notifications_tx + .send(notification) + .map_err(Error::into_internal_error) + }, + agent_client_protocol::on_receive_notification!(), + ); + + let result = Client.v2().connect_with(agent, async move |cx| { + cx.send_request(v2::InitializeRequest::new(ProtocolVersion::V2, + v2::Implementation::new("native-rmcp-client", "1"))).block_task().await?; + let created = Arc::new(std::sync::atomic::AtomicUsize::new(0)); + let factory_calls = created.clone(); + let server = McpServer::::from_rmcp("real-rmcp", move || { + factory_calls.fetch_add(1, std::sync::atomic::Ordering::SeqCst); + Service { + _drop: DropSignal(dropped.clone()), + started: started.clone(), stopped: stopped.clone(), + pending: pending_handlers.clone(), + } + }); + assert_eq!(created.load(std::sync::atomic::Ordering::SeqCst), 0); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)?.start_session().block_task().await?; + let server_id = result_rx.await.map_err(Error::into_internal_error)??; + assert_eq!(created.load(std::sync::atomic::Ordering::SeqCst), 1, + "independent native operations share one application service"); + let acknowledgment = notifications_rx.recv().await.expect("acknowledgment"); + let update = notifications_rx.recv().await.expect("filtered update"); + assert_eq!(acknowledgment.method, "notifications/subscriptions/acknowledged"); + assert_eq!( + acknowledgment.params.as_ref().unwrap()["notifications"]["toolsListChanged"], + json!(true), + "the rmcp subscription must accept the requested notification filter" + ); + assert_eq!(update.method, "notifications/tools/list_changed"); + for notification in [acknowledgment, update] { + assert_eq!(notification.server_id, server_id); + assert_eq!(notification.request_id.0.as_ref(), "listen-1"); + assert_eq!(notification.params.as_ref().unwrap()["_meta"] + ["io.modelcontextprotocol/subscriptionId"], json!("listen-1")); + } + Ok(()) + }).await; + drop_rx.await.map_err(Error::into_internal_error)?; + result + }) + .await + .expect("native ACP/rmcp operation or cleanup timed out") +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn active_rmcp_handler_drops_before_clean_acp_eof_completes() -> Result<(), Error> { + tokio::time::timeout(Duration::from_secs(10), async { + let (handler_started_tx, handler_started_rx) = oneshot::channel(); + let (handler_dropped_tx, mut handler_dropped_rx) = oneshot::channel(); + let pending = Arc::new(Mutex::new(HashMap::from([( + "eof".to_owned(), + (handler_started_tx, handler_dropped_tx), + )]))); + let (peer_stop_tx, peer_stop_rx) = oneshot::channel::<()>(); + let (peer, client) = Channel::duplex(); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx: V2ConnectionTo| { + responder.respond( + v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("rmcp-eof-agent", "1"), + ) + .capabilities( + v2::AgentCapabilities::new().session( + v2::SessionCapabilities::new().mcp( + v2::McpCapabilities::new().acp(v2::McpAcpCapabilities::new()), + ), + ), + ), + ) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected native rmcp declaration"); + }; + let server_id = server.server_id.clone(); + let call_cx = cx.clone(); + cx.spawn(async move { + let mut params = json!({"name": "hang", "arguments": {"probe": "eof"}}); + params["_meta"] = meta("eof"); + let _result = call_cx + .send_request( + v2::MessageMcpRequest::new(server_id, "rmcp-eof", "tools/call") + .params(params.as_object().unwrap().clone()), + ) + .block_task() + .await; + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new("rmcp-eof-session")) + }, + agent_client_protocol::on_receive_request!(), + ); + let peer_task = tokio::spawn(agent.connect_with(peer, async move |_cx| { + let _ = peer_stop_rx.await; + Ok(()) + })); + let client_task = tokio::spawn(Client.v2().connect_with(client, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("rmcp-eof-client", "1"), + )) + .block_task() + .await?; + let pending = pending.clone(); + let server = McpServer::::from_rmcp("rmcp-eof", move || { + let (subscription_started, _) = oneshot::channel(); + let (subscription_stopped, _) = oneshot::channel(); + Service { + _drop: DropSignal(Arc::new(Mutex::new(None))), + started: Arc::new(Mutex::new(Some(subscription_started))), + stopped: Arc::new(Mutex::new(Some(subscription_stopped))), + pending: pending.clone(), + } + }); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)? + .start_session() + .block_task() + .await?; + cx.incoming_closed().await; + Ok(()) + })); + + handler_started_rx + .await + .map_err(Error::into_internal_error)?; + let _ = peer_stop_tx.send(()); + peer_task.await.map_err(Error::into_internal_error)??; + client_task.await.map_err(Error::into_internal_error)??; + assert!( + matches!(handler_dropped_rx.try_recv(), Ok(())), + "rmcp handler future must be destroyed before ACP driver completion" + ); + Ok(()) + }) + .await + .expect("rmcp handler close/join on ACP EOF timed out") +} diff --git a/src/agent-client-protocol-test/src/testy.rs b/src/agent-client-protocol-test/src/testy.rs index 9ae4b604..61ca2e9b 100644 --- a/src/agent-client-protocol-test/src/testy.rs +++ b/src/agent-client-protocol-test/src/testy.rs @@ -1464,15 +1464,26 @@ impl Testy { operation: F, ) -> Result where - F: FnOnce(rmcp::service::RunningService) -> Fut, + F: FnOnce( + rmcp::service::RunningService, + ) -> Fut, Fut: std::future::Future>, { use rmcp::{ - ServiceExt, + ClientLifecycleMode, ClientServiceExt, ServiceExt, + model::{ClientCapabilities, ClientConfig, Implementation, ProtocolVersion}, transport::{ConfigureCommandExt, TokioChildProcess}, }; use tokio::process::Command; + let client_config = || { + ClientConfig::new( + ClientCapabilities::default(), + Implementation::new("testy", env!("CARGO_PKG_VERSION")), + ) + .with_protocol_version(ProtocolVersion::V_2026_07_28) + }; + let mcp_servers = self .get_mcp_servers(session_id) .ok_or_else(|| anyhow::anyhow!("Session not found"))?; @@ -1490,7 +1501,9 @@ impl Testy { match mcp_server { McpServer::Stdio(stdio) => { self.run_until_session_cancelled(session_id, async move { - let mcp_client = () + // Standalone stdio servers may still require initialize; + // native-over-ACP HTTP below uses discover without fallback. + let mcp_client = ClientConfig::default() .serve(TokioChildProcess::new( Command::new(&stdio.command).configure(|cmd| { cmd.args(&stdio.args); @@ -1516,9 +1529,14 @@ impl Testy { .custom_headers(http_headers(&http.headers)?); self.run_until_session_cancelled(session_id, async move { - let mcp_client = - ().serve(StreamableHttpClientTransport::from_config(transport_config)) - .await?; + let mcp_client = client_config() + .serve_with_lifecycle( + StreamableHttpClientTransport::from_config(transport_config), + ClientLifecycleMode::Discover { + preferred_versions: vec![ProtocolVersion::V_2026_07_28], + }, + ) + .await?; operation(mcp_client).await }) diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index fa0324ee..2ad091a2 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -2,6 +2,26 @@ ## [Unreleased] +### Changed (unstable MCP-over-ACP) + +- Target MCP 2026-07-28 with server-addressed `mcp/message` operations and + logical `McpRequestId`s. Remove connect/disconnect and reverse MCP requests; + providers send request-scoped notifications and use ACP cancellation. +- Use reusable native services with owned per-operation execution and cleanup; + retain an explicit connector-backed adapter for factory-based servers. + Expose `request_id()` in attached MCP contexts instead of `connection_id()`. + Preserve standalone serving independently of the unstable ACP feature. +- Validate modern request metadata, restrict discovery to the binding's MCP + revision, and bound native admission, payloads, and transport queues. +- Preserve MCP error codes, omitted/null data, and extensions in a distinct + inner outcome carrier, including connector-backed byte-stream servers. + +### Changed (breaking transport APIs) + +- Raw JSON-RPC responses use `RawJsonRpcResponse` and `RawJsonRpcError`, + preserving error fields without ACP interpretation. Typed ACP consumers + continue receiving `Error`; raw adapters must use the new response type. + ### Added - Add a default-enabled `schemars` feature that forwards JSON Schema support to diff --git a/src/agent-client-protocol/Cargo.toml b/src/agent-client-protocol/Cargo.toml index 4939a6e8..61f79dc3 100644 --- a/src/agent-client-protocol/Cargo.toml +++ b/src/agent-client-protocol/Cargo.toml @@ -62,6 +62,7 @@ wasm_js = ["uuid/js"] [dependencies] agent-client-protocol-schema.workspace = true agent-client-protocol-derive.workspace = true +async-channel.workspace = true futures.workspace = true futures-concurrency.workspace = true rustc-hash.workspace = true diff --git a/src/agent-client-protocol/examples/v2_session_coordination/tests.rs b/src/agent-client-protocol/examples/v2_session_coordination/tests.rs index a8a8aea7..68619f0b 100644 --- a/src/agent-client-protocol/examples/v2_session_coordination/tests.rs +++ b/src/agent-client-protocol/examples/v2_session_coordination/tests.rs @@ -1,7 +1,8 @@ use std::{future::Future, time::Duration}; use agent_client_protocol::{ - Channel, RawJsonRpcMessage, TransportBatch, TransportFrame, schema::v1::RequestId, + BudgetedFrame, Channel, RawJsonRpcMessage, TransportBatch, TransportFrame, + schema::v1::RequestId, }; use serde_json::{Value, json}; @@ -15,7 +16,7 @@ struct Peer(Channel); impl Peer { async fn request(&mut self, method: &str, session: Option<&str>) -> RequestId { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - self.0.rx.next().await + self.0.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected {method}"); }; @@ -34,7 +35,7 @@ impl Peer { fn respond(&self, id: RequestId, result: Result) { self.0 .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .try_send(TransportFrame::Single(RawJsonRpcMessage::response( id, result, ))) .unwrap(); @@ -43,7 +44,7 @@ impl Peer { fn replay_and_respond(&self, id: RequestId, session: &str, text: &str) { self.0 .tx - .unbounded_send(TransportFrame::Batch( + .try_send(TransportFrame::Batch( TransportBatch::from_messages([ update(session, text), RawJsonRpcMessage::response(id, Ok(json!({}))), @@ -83,7 +84,7 @@ impl Peer { while let Some(frame) = self.0.rx.next().await { assert!( matches!( - frame, + frame.frame(), TransportFrame::Single(RawJsonRpcMessage::Notification(_)) ), "unexpected request during shutdown: {frame:?}" @@ -191,7 +192,8 @@ async fn concurrent_loaders_share_one_resume_and_projection() { abandon.await.unwrap(); peer.0 .tx - .unbounded_send(TransportFrame::Single(update(SESSION, "hello"))) + .send_frame(TransportFrame::Single(update(SESSION, "hello"))) + .await .unwrap(); peer.respond(resume, Ok(response)); // A second resume or an early/duplicate close fails this script. @@ -228,13 +230,14 @@ async fn abandoned_resume_is_drained_and_closed_before_fresh_replay() { // Pre-close traffic must also drain before installing a new recipient. peer.0 .tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages([ update(SESSION, "closing"), RawJsonRpcMessage::response(close, Ok(json!({}))), ]) .unwrap(), )) + .await .unwrap(); let fresh = peer.request("session/resume", Some(SESSION)).await; peer.replay_and_respond(fresh, SESSION, "hello"); @@ -277,13 +280,14 @@ async fn delayed_close_blocks_only_its_session() { release_close.await.unwrap(); peer.0 .tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages([ update(OTHER, "+live"), RawJsonRpcMessage::response(close, Ok(json!({}))), ]) .unwrap(), )) + .await .unwrap(); let fresh = peer.request("session/resume", Some(SESSION)).await; peer.replay_and_respond(fresh, SESSION, "fresh"); @@ -370,7 +374,7 @@ async fn disconnect_during_cleanup_fails_waiting_reopen() { // A replacement resume must never have been published. while let Some(frame) = peer.0.rx.next().await { assert!(matches!( - frame, + frame.into_frame(), TransportFrame::Single(RawJsonRpcMessage::Notification(_)) )); } @@ -511,7 +515,7 @@ async fn eof_follows_received_replay_and_response_but_fails_unanswered_loads() { drop(peer.0.tx); while let Some(frame) = peer.0.rx.next().await { assert!(matches!( - frame, + frame.into_frame(), TransportFrame::Single(RawJsonRpcMessage::Notification(_)) )); } diff --git a/src/agent-client-protocol/src/acp_agent.rs b/src/agent-client-protocol/src/acp_agent.rs index 038fb388..9c103b04 100644 --- a/src/agent-client-protocol/src/acp_agent.rs +++ b/src/agent-client-protocol/src/acp_agent.rs @@ -1344,7 +1344,7 @@ mod tests { #[cfg(unix)] async fn reported_descendant_pid( - connection: &mut futures::future::BoxFuture<'static, Result<(), crate::Error>>, + connection: &mut (impl Future> + Unpin), pid_rx: &mut tokio::sync::mpsc::UnboundedReceiver, ) -> rustix::process::Pid { tokio::time::timeout(std::time::Duration::from_secs(5), async { @@ -1417,7 +1417,7 @@ mod tests { Ok(serde_json::json!({ "payload": "x".repeat(4 * 1024 * 1024) })), ); outgoing - .unbounded_send(crate::TransportFrame::Single(response)) + .try_send(crate::TransportFrame::Single(response)) .expect("response should be accepted before the connection starts"); outgoing.close_channel(); diff --git a/src/agent-client-protocol/src/component.rs b/src/agent-client-protocol/src/component.rs index ba8b9297..75b78c4a 100644 --- a/src/agent-client-protocol/src/component.rs +++ b/src/agent-client-protocol/src/component.rs @@ -27,10 +27,61 @@ //! ``` use futures::future::BoxFuture; -use std::{fmt::Debug, future::Future, marker::PhantomData}; +use std::{ + fmt::Debug, + future::Future, + marker::PhantomData, + pin::Pin, + task::{Context, Poll}, +}; use crate::{Channel, Result, role::Role}; +/// Connection work owned by a component, or a passive endpoint with no driver. +/// +/// Both can be awaited, but successful completion of a passive driver says +/// nothing about endpoint lifetime. Bridges must continue copying both halves +/// until they close. An active driver owns the component's completion signal. +pub struct ConnectionDriver(Option>>); + +impl ConnectionDriver { + /// Wrap work that owns a component's connection lifetime. + pub fn new(future: impl Future> + Send + 'static) -> Self { + Self(Some(Box::pin(future))) + } + + /// An endpoint whose I/O is driven elsewhere, such as an existing Channel. + #[must_use] + pub fn passive() -> Self { + Self(None) + } + + /// Whether completion is a no-op rather than an owned lifetime signal. + #[must_use] + pub fn is_passive(&self) -> bool { + self.0.is_none() + } +} + +impl Future for ConnectionDriver { + type Output = Result<()>; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + match self.0.as_mut() { + Some(future) => future.as_mut().poll(cx), + None => Poll::Ready(Ok(())), + } + } +} + +impl Debug for ConnectionDriver { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("ConnectionDriver") + .field("passive", &self.is_passive()) + .finish() + } +} + /// A component that can exchange JSON-RPC messages to an endpoint playing the role `R` /// (e.g., an ACP [`Agent`](`crate::role::acp::Agent`) or an MCP [`Server`](`crate::role::mcp::Server`)). /// @@ -137,7 +188,8 @@ pub trait ConnectTo: Send + 'static { /// /// This method returns: /// - A `Channel` that can be used to communicate with this component - /// - A `BoxFuture` that drives the component's connection logic + /// - A [`ConnectionDriver`] that drives the component's connection logic, + /// or explicitly identifies an endpoint driven elsewhere /// /// The default implementation creates an intermediate channel pair and calls `connect_to` /// on one endpoint while returning the other endpoint for the caller to use. @@ -146,14 +198,14 @@ pub trait ConnectTo: Send + 'static { /// /// # Returns /// - /// A tuple of `(Channel, BoxFuture)` where the channel is for the caller to use + /// A tuple of `(Channel, ConnectionDriver)` where the channel is for the caller to use /// and the future must be polled to drive the connection. - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<()>>) + fn into_channel_and_future(self) -> (Channel, ConnectionDriver) where Self: Sized, { let (channel_a, channel_b) = Channel::duplex(); - let future = Box::pin(self.connect_to(channel_b)); + let future = ConnectionDriver::new(self.connect_to(channel_b)); (channel_a, future) } } @@ -171,8 +223,7 @@ trait ErasedConnectTo: Send { client: Box>, ) -> BoxFuture<'static, Result<()>>; - fn into_channel_and_future_erased(self: Box) - -> (Channel, BoxFuture<'static, Result<()>>); + fn into_channel_and_future_erased(self: Box) -> (Channel, ConnectionDriver); } /// Blanket implementation: any `ConnectTo` can be type-erased. @@ -195,9 +246,7 @@ impl, R: Role> ErasedConnectTo for C { }) } - fn into_channel_and_future_erased( - self: Box, - ) -> (Channel, BoxFuture<'static, Result<()>>) { + fn into_channel_and_future_erased(self: Box) -> (Channel, ConnectionDriver) { (*self).into_channel_and_future() } } @@ -251,7 +300,7 @@ impl ConnectTo for DynConnectTo { .await } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<()>>) { + fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { self.inner.into_channel_and_future_erased() } } diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 699ae9cf..76671ef8 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -2,7 +2,7 @@ use agent_client_protocol_schema::v1::{ JsonRpcMessage as VersionedJsonRpcMessage, Notification as RpcNotification, - Request as RpcRequest, RequestId, Response as RpcResponse, SessionId, + Request as RpcRequest, RequestId, SessionId, }; // Types re-exported from crate root @@ -18,19 +18,22 @@ use std::sync::{ Arc, Mutex, Weak, atomic::{AtomicBool, Ordering}, }; +use std::task::{Context, Poll, Waker}; use uuid::Uuid; use futures::FutureExt; -use futures::channel::{mpsc, oneshot}; +use futures::channel::oneshot; use futures::future::{self, BoxFuture, Either}; -use futures::{AsyncRead, AsyncWrite, StreamExt}; +use futures::{AsyncRead, AsyncWrite, Sink, SinkExt, StreamExt}; +mod admission; pub(crate) mod close; mod dynamic_handler; pub(crate) mod handlers; mod incoming_actor; mod outgoing_actor; mod protocol_compat; +mod raw_error; pub(crate) mod run; mod task_actor; mod transport_actor; @@ -43,6 +46,7 @@ use crate::jsonrpc::handlers::{ChainedHandler, NamedHandler}; use crate::jsonrpc::handlers::{MessageHandler, NotificationHandler, RequestHandler}; use crate::jsonrpc::outgoing_actor::{OutgoingMessageTx, send_raw_message}; use crate::jsonrpc::protocol_compat::{ProtocolCompat, ProtocolMode}; +pub use crate::jsonrpc::raw_error::{RawJsonRpcError, RawJsonRpcResponse}; use crate::jsonrpc::run::SpawnedRun; use crate::jsonrpc::run::{ChainRun, NullRun, RunWithConnectionTo}; use crate::jsonrpc::task_actor::{Task, TaskTx}; @@ -62,8 +66,8 @@ pub enum RawJsonRpcMessage { Request(RpcRequest), /// A JSON-RPC notification without a response. Notification(RpcNotification), - /// A JSON-RPC response to a prior request. - Response(RpcResponse), + /// A response with an opaque result or transport-neutral error object. + Response(RawJsonRpcResponse), } /// A JSON-RPC frame exchanged between protocol components and transports. @@ -87,6 +91,32 @@ pub enum TransportFrame { Batch(TransportBatch), } +/// Finite transport and runtime admission limits. The byte budget is shared +/// across both directions of one in-memory duplex. +#[derive(Clone, Copy, Debug)] +pub struct ConnectionLimits { + /// Maximum UTF-8 bytes in one JSON-RPC frame. + pub max_frame_bytes: usize, + /// Shared serialized-payload budget, including queued frames and runtime + /// messages. One maximum frame's worth is reserved for responses/cancellation. + pub max_queued_bytes: usize, + /// Per-queue item limit and runtime admission limit for pending requests, + /// total live tasks (running plus waiting), dynamic handlers, and deferred + /// dispatch. Values below one are treated as one. Byte capacity is enforced + /// separately. + pub max_queued_frames: usize, +} + +impl Default for ConnectionLimits { + fn default() -> Self { + Self { + max_frame_bytes: transport_actor::MAX_FRAME_BYTES, + max_queued_bytes: 64 * 1024 * 1024, + max_queued_frames: admission::QUEUE_CAPACITY, + } + } +} + /// A structurally non-empty JSON-RPC batch retained across framed relays. #[derive(Clone, Debug)] pub struct TransportBatch { @@ -236,6 +266,48 @@ impl Serialize for TransportBatch { } impl TransportFrame { + fn is_control(&self) -> bool { + fn message_is_control(message: &RawJsonRpcMessage) -> bool { + match message { + RawJsonRpcMessage::Response(_) => true, + RawJsonRpcMessage::Notification(notification) => { + if matches!( + notification.method.as_ref(), + "$/cancel_request" | "$/cancelRequest" + ) { + return true; + } + if !crate::schema::SuccessorMessage::::matches_method( + ¬ification.method, + ) { + return false; + } + let Some(RawJsonRpcParams::Object(envelope)) = ¬ification.params else { + return false; + }; + let Some(method) = envelope.get("method").and_then(serde_json::Value::as_str) + else { + return false; + }; + let (method, _) = peel_successor_envelopes( + method, + envelope.get("params").unwrap_or(&serde_json::Value::Null), + ); + matches!(method, "$/cancel_request" | "$/cancelRequest") + } + RawJsonRpcMessage::Request(_) => false, + } + } + match self { + Self::Single(message) => message_is_control(message), + Self::Batch(batch) => batch.entries().all(|entry| match entry { + TransportBatchEntry::Message(message) => message_is_control(message), + TransportBatchEntry::Malformed { .. } => false, + }), + Self::Malformed { .. } => false, + } + } + fn inspect_messages( &self, observer: &mut impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error>, @@ -338,19 +410,25 @@ impl RawJsonRpcMessage { })) } - /// Build a raw JSON-RPC response message. + /// Build a raw response from an ACP result. + /// + /// For other protocols, construct [`RawJsonRpcResponse`] directly so error + /// codes and fields are not first interpreted as ACP errors. #[must_use] pub fn response(id: RequestId, response: Result) -> Self { - Self::Response(RpcResponse::new(id, response)) + Self::Response(RawJsonRpcResponse::new( + id, + response.map_err(|error| Box::new(error.into())), + )) } /// The response id, if this is a response. #[must_use] pub fn response_id(&self) -> Option<&RequestId> { match self { - Self::Response(RpcResponse::Result { id, .. } | RpcResponse::Error { id, .. }) => { - Some(id) - } + Self::Response( + RawJsonRpcResponse::Result { id, .. } | RawJsonRpcResponse::Error { id, .. }, + ) => Some(id), Self::Request(_) | Self::Notification(_) => None, } } @@ -407,11 +485,10 @@ impl<'de> Deserialize<'de> for RawJsonRpcMessage { Ok(Self::Notification(notification)) } } else if !has_method && has_id && has_result != has_error { - let response = serde_json::from_value::< - VersionedJsonRpcMessage>, - >(value) - .map_err(serde::de::Error::custom)? - .into_inner(); + let response = + serde_json::from_value::>(value) + .map_err(serde::de::Error::custom)? + .into_inner(); Ok(Self::Response(response)) } else { Err(serde::de::Error::custom("invalid JSON-RPC message")) @@ -1898,15 +1975,22 @@ impl< context: _, } = self; - let (outgoing_tx, outgoing_rx) = mpsc::unbounded(); - let (new_task_tx, new_task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, dynamic_handler_rx) = mpsc::unbounded(); - let pending_replies = PendingReplies::default(); - // Convert transport into server - this returns a channel for us to use // and a future that runs the transport. let transport_component = crate::DynConnectTo::new(transport); let (transport_channel, transport_future) = transport_component.into_channel_and_future(); + let limits = transport_channel.tx.admission().limits(); + let (outgoing_tx, outgoing_rx) = admission::budgeted_channel( + transport_channel.tx.admission(), + OutgoingMessage::charged_bytes, + OutgoingMessage::with_permit, + OutgoingMessage::is_control, + OutgoingMessage::is_urgent, + ); + let (new_task_tx, new_task_rx) = task_actor::task_channel(limits.max_queued_frames); + let (dynamic_handler_tx, dynamic_handler_rx) = + admission::channel_with_capacity(limits.max_queued_frames); + let pending_replies = PendingReplies::with_capacity(limits.max_queued_frames); let (transport_completion_tx, transport_completion_rx) = oneshot::channel(); let transport_completion = transport_completion_rx .map(|result| { @@ -1928,11 +2012,11 @@ impl< pending_replies.registrar(), protocol_mode, ); - let spawn_result = connection.spawn(async move { + let transport_driver = async move { let result = transport_future.await; drop(transport_completion_tx.send(result.clone())); result - }); + }; // Destructure the channel endpoints let Channel { @@ -1945,8 +2029,6 @@ impl< let future = crate::util::instrument_with_connection_name(name, { let connection = connection.clone(); async move { - let () = spawn_result?; - let background = async { let incoming = incoming_actor::incoming_protocol_actor( me.counterpart(), @@ -1957,20 +2039,13 @@ impl< incoming_actor::IncomingHandlers::new(handler, on_close), protocol_compat.clone(), ); - let other_actors = async { - futures::try_join!( - // Protocol layer: OutgoingMessage -> RawJsonRpcMessage - outgoing_actor::outgoing_protocol_actor( - outgoing_rx, - pending_replies, - transport_outgoing_tx, - protocol_compat, - ), - task_actor::task_actor(new_task_rx, &connection), - runner.run_with_connection_to(connection.clone()), - )?; - Ok(()) - }; + let other_actors = outgoing_actor::outgoing_protocol_actor( + outgoing_rx, + pending_replies, + transport_outgoing_tx, + protocol_compat, + connection.incoming_closed.clone(), + ); // EOF can wake a pending request consumer, which may make // the task actor fail while close callbacks are running. @@ -1984,10 +2059,45 @@ impl< .await }; - run_until_connection_close( - background, - main_fn(connection.clone()), + let lifecycle = run_until_connection_close( + async { + let result = background.await; + connection.incoming_closed.request_shutdown(); + result + }, + async { + let result = main_fn(connection.clone()).await; + connection.incoming_closed.request_shutdown(); + result + }, connection.incoming_closed.clone(), + ); + // Only native operation supervisors are joined at shutdown. + // Ordinary spawned tasks and user runners remain disposable. + crate::util::run_until( + finish_actor_error( + runner.run_with_connection_to(connection.clone()), + &connection, + ), + crate::util::run_until( + finish_actor_error(transport_driver, &connection), + crate::util::run_until( + finish_actor_error( + task_actor::task_actor( + new_task_rx, + &connection, + limits.max_queued_frames, + ), + &connection, + ), + async { + let result = lifecycle.await; + connection.incoming_closed.request_shutdown(); + connection.wait_protected_operations().await; + result + }, + ), + ), ) .await } @@ -1997,6 +2107,23 @@ impl< } } +/// An EOF may wake a task that reports an error while on-close callbacks are +/// still running. Preserve that error without dropping the callback future. +async fn finish_actor_error( + future: impl Future>, + connection: &ConnectionTo, +) -> Result<(), crate::Error> { + let result = future.await; + if result.is_err() { + connection.incoming_closed.request_shutdown(); + if connection.incoming_closed.is_closing() { + connection.incoming_closed.closed().await; + } + connection.wait_protected_operations().await; + } + result +} + #[cfg(feature = "unstable_mcp_over_acp")] impl< Host: Role, @@ -2098,6 +2225,8 @@ pub(crate) struct ResponsePayload { /// the dispatch loop; ordinary blocking consumers, local error paths, and /// responses routed later do not. pub(crate) ack_tx: Option>, + /// Admission remains with an SDK-owned result until it is consumed or dropped. + retained_bytes: Option, } type ResponseRouteHook = @@ -2141,7 +2270,7 @@ impl std::fmt::Debug for ResponsePayload { f.debug_struct("ResponsePayload") .field("result", &self.result) .field("ack_tx", &self.ack_tx.as_ref().map(|_| "...")) - .finish() + .finish_non_exhaustive() } } @@ -2162,6 +2291,8 @@ impl ResponseOrdering { struct PendingReply { method: String, + /// The method and map key outlive the outgoing frame. + metadata_bytes: Option, role_id: RoleId, sender: oneshot::Sender, cancellation_disarm: SentRequestCancellationDisarm, @@ -2177,6 +2308,7 @@ impl PendingReply { .send(ResponsePayload { result: Err(error), ack_tx: None, + retained_bytes: self.metadata_bytes, }) .is_err() { @@ -2190,10 +2322,20 @@ impl PendingReply { } } -#[derive(Default)] struct PendingRepliesInner { incoming_closed: bool, replies: HashMap, + max_pending: usize, +} + +impl Default for PendingRepliesInner { + fn default() -> Self { + Self { + incoming_closed: false, + replies: HashMap::new(), + max_pending: admission::QUEUE_CAPACITY, + } + } } #[derive(Clone, Default)] @@ -2202,6 +2344,15 @@ struct PendingReplies { } impl PendingReplies { + fn with_capacity(max_pending: usize) -> Self { + Self { + inner: Arc::new(Mutex::new(PendingRepliesInner { + max_pending: max_pending.max(1), + ..PendingRepliesInner::default() + })), + } + } + fn registrar(&self) -> PendingRepliesRegistrar { PendingRepliesRegistrar { inner: Arc::downgrade(&self.inner), @@ -2224,6 +2375,36 @@ impl PendingReplies { .remove(id) } + fn mark_published(&self, id: &RequestId) -> bool { + let inner = self.inner.lock().expect("pending replies mutex poisoned"); + let Some(reply) = inner.replies.get(id) else { + return false; + }; + reply + .cancellation_disarm + .published + .store(true, Ordering::Release); + true + } + + /// Cancellation may bypass queued work, but must never reach the peer + /// before a request that we subsequently publish. Settle that case locally. + fn cancel_unpublished(&self, id: &RequestId) -> bool { + let reply = { + let mut inner = self.inner.lock().expect("pending replies mutex poisoned"); + if inner + .replies + .get(id) + .is_none_or(|reply| reply.cancellation_disarm.published.load(Ordering::Acquire)) + { + return false; + } + inner.replies.remove(id).expect("pending reply checked") + }; + reply.fail(crate::Error::request_cancelled()); + true + } + /// Atomically reject new subscriptions and fail every existing one. fn close_incoming(&self) -> usize { let replies = { @@ -2274,6 +2455,12 @@ impl PendingRepliesRegistrar { let mut inner = inner.lock().expect("pending replies mutex poisoned"); if inner.incoming_closed { Err(reply) + } else if !inner.replies.contains_key(&id) && inner.replies.len() >= inner.max_pending { + drop(inner); + reply.fail(crate::util::internal_error( + "pending request capacity exceeded", + )); + return false; } else { Ok(inner.replies.insert(id, reply)) } @@ -2304,6 +2491,17 @@ impl PendingRepliesRegistrar { .replies .remove(id) } + + fn discard_abandoned(&self, id: &RequestId) -> Option { + let inner = self.inner.upgrade()?; + let mut inner = inner.lock().expect("pending replies mutex poisoned"); + // Framework response hooks own cleanup even when their consumer drops. + // Keep their bounded registration until the reply arrives or EOF fails it. + if inner.replies.get(id)?.response_route_hook.is_some() { + return None; + } + inner.replies.remove(id) + } } impl Debug for PendingRepliesRegistrar { @@ -2653,6 +2851,14 @@ fn peel_successor_envelopes<'message>( (method, params) } +fn outgoing_cancellation_id(message: &UntypedMessage) -> Option { + let (method, params) = peel_successor_envelopes(&message.method, &message.params); + if !matches!(method, "$/cancel_request" | "$/cancelRequest") { + return None; + } + serde_json::from_value(params.get("requestId")?.clone()).ok() +} + /// Whether a notification is a `$/cancel_request`, even when it is still /// wrapped in `_proxy/successor` envelopes. /// @@ -2718,6 +2924,7 @@ impl ResponseDestination { remaining: slot_count, responses: (0..slot_count).map(|_| None).collect(), abandoned: (0..slot_count).map(|_| None).collect(), + permits: (0..slot_count).map(|_| None).collect(), active_handler_attempts: (0..slot_count).map(|_| 0).collect(), dispatch_complete: false, emitted: false, @@ -2737,17 +2944,29 @@ impl ResponseDestination { ) } - fn complete(self, response: RawJsonRpcMessage) -> Option { + fn complete_admitted( + self, + response: RawJsonRpcMessage, + permit: Option, + ) -> Option<(TransportFrame, Option)> { match self { - Self::Individual(slot) => slot.complete(response), - Self::Batch(slot) => slot.complete(response).map(batch_response_frame), + Self::Individual(slot) => slot.complete(response).map(|frame| (frame, permit)), + Self::Batch(slot) => slot + .complete_admitted(response, permit) + .map(batch_response_frame_admitted), } } - fn abandon(self, fallback: RawJsonRpcMessage) -> Option { + fn abandon_admitted( + self, + fallback: RawJsonRpcMessage, + permit: Option, + ) -> Option<(TransportFrame, Option)> { match self { Self::Individual(_) => None, - Self::Batch(slot) => slot.abandon(fallback).map(batch_response_frame), + Self::Batch(slot) => slot + .abandon_admitted(fallback, permit) + .map(batch_response_frame_admitted), } } @@ -2769,10 +2988,19 @@ impl ResponseDestination { }) } - fn finish_handler_attempt(self) -> Option { + fn finish_handler_attempt_admitted( + self, + permit: Option, + ) -> Option<(TransportFrame, Option)> { match self { Self::Individual(_) => None, - Self::Batch(slot) => slot.finish_handler_attempt().map(batch_response_frame), + Self::Batch(slot) => slot.finish_handler_attempt().map(|ready| { + let (frame, mut charge) = batch_response_frame_admitted(ready); + if let (Some(charge), Some(permit)) = (&mut charge, permit) { + charge.join(permit); + } + (frame, charge) + }), } } } @@ -2800,6 +3028,20 @@ fn batch_response_frame(responses: Vec) -> TransportFrame { ) } +fn batch_response_frame_admitted(ready: BatchReady) -> (TransportFrame, Option) { + let mut permits = ready.permits.into_iter(); + let mut charge = permits.next(); + for permit in permits { + charge.as_mut().expect("first permit exists").join(permit); + } + (batch_response_frame(ready.responses), charge) +} + +struct BatchReady { + responses: Vec, + permits: Vec, +} + #[derive(Clone)] struct BatchDispatchCompletion { state: Arc>, @@ -2814,7 +3056,10 @@ impl std::fmt::Debug for BatchDispatchCompletion { } impl BatchDispatchCompletion { - fn complete(self) -> Option { + fn complete_admitted( + self, + permit: Option, + ) -> Option<(TransportFrame, Option)> { let mut state = self .state .lock() @@ -2827,7 +3072,13 @@ impl BatchDispatchCompletion { for index in 0..state.responses.len() { promote_abandoned_response(&mut state, index); } - take_completed_batch(&mut state).map(batch_response_frame) + take_completed_batch(&mut state).map(|ready| { + let (frame, mut charge) = batch_response_frame_admitted(ready); + if let (Some(charge), Some(permit)) = (&mut charge, permit) { + charge.join(permit); + } + (frame, charge) + }) } } @@ -2841,14 +3092,14 @@ fn promote_abandoned_response(state: &mut BatchResponseState, index: usize) { } } -fn take_completed_batch(state: &mut BatchResponseState) -> Option> { +fn take_completed_batch(state: &mut BatchResponseState) -> Option { if !state.dispatch_complete || state.remaining != 0 || state.emitted { return None; } state.emitted = true; - Some( - state + Some(BatchReady { + responses: state .responses .iter_mut() .map(|response| { @@ -2857,7 +3108,8 @@ fn take_completed_batch(state: &mut BatchResponseState) -> Option Option> { + fn finish_handler_attempt(self) -> Option { let mut state = self .state .lock() @@ -2898,7 +3150,11 @@ impl BatchResponseSlot { take_completed_batch(&mut state) } - fn complete(self, response: RawJsonRpcMessage) -> Option> { + fn complete_admitted( + self, + response: RawJsonRpcMessage, + permit: Option, + ) -> Option { let mut state = self .state .lock() @@ -2924,11 +3180,16 @@ impl BatchResponseSlot { state.abandoned[self.index] = None; state.responses[self.index] = Some(response); + state.permits[self.index] = permit; state.remaining -= 1; take_completed_batch(&mut state) } - fn abandon(self, fallback: RawJsonRpcMessage) -> Option> { + fn abandon_admitted( + self, + fallback: RawJsonRpcMessage, + permit: Option, + ) -> Option { let mut state = self .state .lock() @@ -2950,6 +3211,7 @@ impl BatchResponseSlot { } else { state.abandoned[self.index] = Some(fallback); } + state.permits[self.index] = permit; take_completed_batch(&mut state) } } @@ -2958,6 +3220,7 @@ struct BatchResponseState { remaining: usize, responses: Vec>, abandoned: Vec>, + permits: Vec>, active_handler_attempts: Vec, dispatch_complete: bool, emitted: bool, @@ -2995,6 +3258,8 @@ struct ResponseReplyTarget { sender: Arc>>>, ordering: ResponseOrdering, dispatch: ResponseDispatch, + /// Keep the original frame admitted while a handler defers routing. + frame_bytes: Option, } impl ResponseReplyTarget { @@ -3013,8 +3278,39 @@ impl ResponseReplyTarget { return; }; + // A transformed result may be larger than the wire response. Each + // result (including each member of a batch) therefore needs its own + // charge; cloning the batch's frame permit does not charge each result. + // Never wait here: this router may hold the only permit whose release + // would make room. On rejection deliver a bounded error instead. + let (result, retained_bytes) = if let Some(frame) = self.frame_bytes { + let bytes = match &result { + Ok(value) => serde_json::to_vec(value).map(|json| json.len()), + Err(error) => serde_json::to_vec(error).map(|json| json.len()), + }; + match bytes.ok().and_then(|bytes| { + FrameAdmission(frame.inner.budget.clone()).try_reserve_bytes(bytes, true) + }) { + Some(permit) => (result, Some(permit)), + None => ( + Err(crate::util::internal_error( + "retained response byte capacity exceeded", + )), + Some(frame), + ), + } + } else { + (result, None) + }; let ack_tx = self.dispatch.acknowledgment(&self.ordering); - if sender.send(ResponsePayload { result, ack_tx }).is_err() { + if sender + .send(ResponsePayload { + result, + ack_tx, + retained_bytes, + }) + .is_err() + { tracing::debug!( method = %self.method, id = ?self.id, @@ -3064,7 +3360,7 @@ impl ResponseDispatch { enum HandlerErrorTarget { Request(RequestReplyTarget), - Response(ResponseReplyTarget), + Response(Box), } impl HandlerErrorTarget { @@ -3081,6 +3377,13 @@ impl HandlerErrorTarget { #[derive(Debug)] enum OutgoingMessage { + /// Retain application admission across queueing, readiness, conversion, and + /// transport publication. Legacy test-only queues can still carry bare messages. + Admitted { + message: Box, + permit: FramePermit, + }, + /// Close the outgoing application queue and acknowledge after every /// already-accepted message has entered the raw transport queue. CloseAfterDraining { done: oneshot::Sender<()> }, @@ -3147,6 +3450,99 @@ enum OutgoingMessage { }, } +impl OutgoingMessage { + fn charged_bytes(&self) -> Result { + // Include space for the JSON-RPC envelope and request ID. A transformed + // frame that exceeds this estimate must grow the *same* permit, never + // await an independent reservation while retaining the first. + const ENVELOPE: usize = 64; + let bytes = match self { + Self::Admitted { message, .. } => return message.charged_bytes(), + Self::Request { + id, + method, + untyped, + .. + } => { + serde_json::to_vec(&untyped.params) + .map_err(crate::Error::into_internal_error)? + .len() + + untyped.method.len() + + method.len() + + serde_json::to_vec(id) + .map_err(crate::Error::into_internal_error)? + .len() + } + Self::Notification { untyped } => { + serde_json::to_vec(&untyped.params) + .map_err(crate::Error::into_internal_error)? + .len() + + untyped.method.len() + } + Self::Response { + id, + method, + response, + .. + } => { + serde_json::to_vec(response) + .map_err(crate::Error::into_internal_error)? + .len() + + method.len() + + serde_json::to_vec(id) + .map_err(crate::Error::into_internal_error)? + .len() + } + Self::UncorrelatedErrorResponse { error, .. } => serde_json::to_vec(error) + .map_err(crate::Error::into_internal_error)? + .len(), + Self::AbandonedBatchResponse { id, method, .. } => { + method.len() + + serde_json::to_vec(id) + .map_err(crate::Error::into_internal_error)? + .len() + } + Self::CloseAfterDraining { .. } + | Self::BatchDispatchComplete { .. } + | Self::BatchHandlerAttemptComplete { .. } => 0, + }; + Ok(if bytes == 0 { + 1 + } else { + bytes.saturating_add(ENVELOPE) + }) + } + + fn is_control(&self) -> bool { + match self { + Self::Admitted { message, .. } => message.is_control(), + Self::Notification { untyped } => outgoing_cancellation_id(untyped).is_some(), + Self::CloseAfterDraining { .. } + | Self::BatchDispatchComplete { .. } + | Self::BatchHandlerAttemptComplete { .. } + | Self::Response { .. } + | Self::UncorrelatedErrorResponse { .. } + | Self::AbandonedBatchResponse { .. } => true, + Self::Request { .. } => false, + } + } + + fn is_urgent(&self) -> bool { + match self { + Self::Admitted { message, .. } => message.is_urgent(), + Self::Notification { untyped } => outgoing_cancellation_id(untyped).is_some(), + _ => false, + } + } + + fn with_permit(self, permit: FramePermit) -> Self { + Self::Admitted { + message: Box::new(self), + permit, + } + } +} + /// Return type from JrHandler; indicates whether the request was handled or not. #[must_use] #[derive(Debug)] @@ -3229,6 +3625,11 @@ impl V2ConnectionTo { self.inner.incoming_closed().await; } + /// Wait for EOF or connection termination, before close callbacks run. + pub async fn shutdown_requested(&self) { + self.inner.shutdown_requested().await; + } + /// Return whether clean incoming-EOF processing has completed. #[must_use] pub fn is_incoming_closed(&self) -> bool { @@ -3337,6 +3738,17 @@ impl V2ConnectionTo { self.inner.send_notification(notification) } + /// Await outbound capacity outside the dispatch loop. + pub async fn send_notification_async( + &self, + notification: N, + ) -> Result<(), crate::Error> + where + Counterpart: HasPeer, + { + self.inner.send_notification_async(notification).await + } + /// Send an outgoing notification to a specific peer. pub fn send_notification_to( &self, @@ -3349,6 +3761,20 @@ impl V2ConnectionTo { self.inner.send_notification_to(peer, notification) } + /// Await outbound capacity outside the dispatch loop. + pub async fn send_notification_to_async( + &self, + peer: Peer, + notification: N, + ) -> Result<(), crate::Error> + where + Counterpart: HasPeer, + { + self.inner + .send_notification_to_async(peer, notification) + .await + } + /// Send a `$/cancel_request` notification to the default counterpart peer. pub fn send_cancel_request( &self, @@ -3431,7 +3857,7 @@ pub struct ConnectionTo { counterpart: Counterpart, message_tx: OutgoingMessageTx, task_tx: TaskTx, - dynamic_handler_tx: mpsc::UnboundedSender>, + dynamic_handler_tx: admission::Sender>, transport_completion: SharedTransportCompletion, pending_replies: PendingRepliesRegistrar, #[cfg_attr( @@ -3443,10 +3869,26 @@ pub struct ConnectionTo { )] protocol_mode: ProtocolMode, incoming_closed: IncomingClosed, + protected_operations: Arc>, } type SharedTransportCompletion = future::Shared>>; +#[derive(Default)] +struct ProtectedOperations { + pending: Vec>, + joining: Option>>, +} + +impl Debug for ProtectedOperations { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter + .debug_struct("ProtectedOperations") + .field("pending", &self.pending.len()) + .finish_non_exhaustive() + } +} + #[derive(Clone)] struct IncomingClosed { state: Arc, @@ -3457,23 +3899,45 @@ struct IncomingClosedState { closed: AtomicBool, signal_tx: Mutex>>, signal_rx: future::Shared>, + shutdown_tx: Mutex>>, + shutdown_rx: future::Shared>, } impl IncomingClosed { fn new() -> Self { let (signal_tx, signal_rx) = oneshot::channel(); + let (shutdown_tx, shutdown_rx) = oneshot::channel(); Self { state: Arc::new(IncomingClosedState { closing: AtomicBool::new(false), closed: AtomicBool::new(false), signal_tx: Mutex::new(Some(signal_tx)), signal_rx: signal_rx.map(|_| ()).boxed().shared(), + shutdown_tx: Mutex::new(Some(shutdown_tx)), + shutdown_rx: shutdown_rx.map(|_| ()).boxed().shared(), }), } } fn begin_close(&self) { self.state.closing.store(true, Ordering::Release); + self.request_shutdown(); + } + + fn request_shutdown(&self) { + if let Some(tx) = self + .state + .shutdown_tx + .lock() + .expect("shutdown mutex poisoned") + .take() + { + let _ = tx.send(()); + } + } + + async fn shutdown_requested(&self) { + self.state.shutdown_rx.clone().await; } fn finish_close(&self) { @@ -3583,9 +4047,9 @@ fn run_until_connection_close( impl ConnectionTo { fn new( counterpart: Counterpart, - message_tx: mpsc::UnboundedSender, - task_tx: mpsc::UnboundedSender, - dynamic_handler_tx: mpsc::UnboundedSender>, + message_tx: OutgoingMessageTx, + task_tx: TaskTx, + dynamic_handler_tx: admission::Sender>, transport_completion: SharedTransportCompletion, pending_replies: PendingRepliesRegistrar, protocol_mode: ProtocolMode, @@ -3599,7 +4063,65 @@ impl ConnectionTo { pending_replies, protocol_mode, incoming_closed: IncomingClosed::new(), + protected_operations: Arc::default(), + } + } + + #[cfg(feature = "unstable_mcp_over_acp")] + #[track_caller] + pub(crate) fn spawn_protected( + &self, + task: impl IntoFuture, IntoFuture: Send + 'static>, + ) -> Result<(), crate::Error> { + let (done_tx, done_rx) = oneshot::channel(); + let task = task.into_future(); + let mut state = self + .protected_operations + .lock() + .expect("protected operations poisoned"); + if state.joining.is_some() { + return Err(crate::Error::request_cancelled()); } + // Completed acknowledgments must not accumulate for the connection's + // entire lifetime. With completed entries reaped at each admission, + // this registry is bounded by the shared live-task limit. + state + .pending + .retain_mut(|done| matches!(done.try_recv(), Ok(None))); + self.spawn(async move { + let result = task.await; + let _ = done_tx.send(()); + result + })?; + state.pending.push(done_rx); + Ok(()) + } + + pub(crate) async fn wait_protected_operations(&self) { + let joining = { + let mut state = self + .protected_operations + .lock() + .expect("protected operations poisoned"); + if state.joining.is_none() { + let operations = std::mem::take(&mut state.pending); + state.joining = Some( + async move { + for operation in operations { + let _ = operation.await; + } + } + .boxed() + .shared(), + ); + } + state.joining.as_ref().expect("join initialized").clone() + }; + joining.await; + } + + pub(crate) fn request_shutdown(&self) { + self.incoming_closed.request_shutdown(); } #[cfg(feature = "unstable_protocol_v2")] @@ -3624,6 +4146,12 @@ impl ConnectionTo { self.incoming_closed.closed().await; } + /// Resolves on transport EOF or local completion, before close callbacks + /// or outgoing drain. Cancel connection-owned work when this fires. + pub async fn shutdown_requested(&self) { + self.incoming_closed.shutdown_requested().await; + } + /// Return whether clean incoming-EOF processing has completed. /// /// This remains `false` while [`Builder::on_close`] callbacks are running. @@ -3636,10 +4164,11 @@ impl ConnectionTo { /// the protocol actor, and wait for the transport sink to finish them. async fn drain_outgoing(&self) -> Result<(), crate::Error> { let (done_tx, done_rx) = oneshot::channel(); - let marker_result = send_raw_message( - &self.message_tx, - OutgoingMessage::CloseAfterDraining { done: done_tx }, - ); + let marker_result = self + .message_tx + .send(OutgoingMessage::CloseAfterDraining { done: done_tx }) + .await + .map_err(crate::util::internal_error); let marker_result = match marker_result { Ok(()) => done_rx.await.map_err(|error| { crate::util::internal_error(format!( @@ -4094,13 +4623,18 @@ impl ConnectionTo { } let role_id = peer.role_id(); let remote_style = self.counterpart.remote_style(peer); - let cancellation = - SentRequestCancellation::new(self.message_tx.clone(), remote_style, id.clone()); + let cancellation = SentRequestCancellation::new( + self.message_tx.clone(), + self.pending_replies.clone(), + remote_style, + id.clone(), + ); if self.is_incoming_closing() { cancellation.disarm(); drop(response_tx.send(ResponsePayload { result: Err(incoming_transport_closed_error(&method)), ack_tx: None, + retained_bytes: None, })); return SentRequest::new( id, @@ -4115,12 +4649,47 @@ impl ConnectionTo { match request.to_untyped_message() { Ok(untyped) => { + // The queue's frame charge is released after transport publication, + // but the pending map retains its own copies of the method and ID. + // Charge those strings (plus a fixed entry allowance) separately. + let metadata_bytes = self.message_tx.byte_admission().and_then(|budget| { + budget.try_reserve_bytes( + method + .len() + .saturating_add(match &id { + RequestId::Str(value) => value.len(), + _ => 32, + }) + .saturating_add(64), + true, + ) + }); + if self.message_tx.byte_admission().is_some() && metadata_bytes.is_none() { + cancellation.disarm(); + drop(response_tx.send(ResponsePayload { + result: Err(crate::util::internal_error( + "pending request metadata byte capacity exceeded", + )), + ack_tx: None, + retained_bytes: None, + })); + return SentRequest::new( + id, + method.clone(), + self.task_tx.clone(), + response_rx, + cancellation, + response_ordering, + ) + .map(move |json| ::from_value(&method, json)); + } // Register before enqueueing so incoming EOF can fail every // observable request before close callbacks begin. The // outgoing actor checks that the registration still exists // before sending the request. let pending_reply = PendingReply { method: method.clone(), + metadata_bytes, role_id, sender: response_tx, cancellation_disarm: cancellation.disarm_handle(), @@ -4142,10 +4711,9 @@ impl ConnectionTo { if let Err(error) = self.message_tx.unbounded_send(message) { cancellation.disarm(); - - let OutgoingMessage::Request { id, method, .. } = error.into_inner() else { - unreachable!(); - }; + // A rejected queue item may be wrapped in Admitted. + // Drop it to release its admission before failing the waiter. + drop(error.into_inner()); if let Some(pending_reply) = self.pending_replies.remove(&id) { if self.is_incoming_closing() { @@ -4169,6 +4737,7 @@ impl ConnectionTo { "failed to create untyped request for `{method}`: {err}" ))), ack_tx: None, + retained_bytes: None, }) .unwrap(); } @@ -4212,6 +4781,18 @@ impl ConnectionTo { self.send_notification_to(self.counterpart.clone(), notification) } + /// Await outbound capacity for a producer outside ordered dispatch. + pub async fn send_notification_async( + &self, + notification: N, + ) -> Result<(), crate::Error> + where + Counterpart: HasPeer, + { + self.send_notification_to_async(self.counterpart.clone(), notification) + .await + } + /// Send an outgoing notification to a specific peer (no reply expected). /// /// The message will be transformed according to the [`HasPeer`](crate::role::HasPeer) @@ -4246,6 +4827,25 @@ impl ConnectionTo { ) } + /// Await outbound capacity for a producer outside ordered dispatch. + pub async fn send_notification_to_async( + &self, + peer: Peer, + notification: N, + ) -> Result<(), crate::Error> + where + Counterpart: HasPeer, + { + let remote_style = self.counterpart.remote_style(peer); + let transformed = remote_style.transform_outgoing_message(notification)?; + self.message_tx + .send(OutgoingMessage::Notification { + untyped: transformed, + }) + .await + .map_err(crate::util::internal_error) + } + /// Send a `$/cancel_request` notification for an arbitrary request ID to /// the default counterpart peer. /// @@ -4722,7 +5322,7 @@ pub struct ResponseRouter { send_fn: Box) -> Result<(), crate::Error> + Send>, /// Shared route used to deliver a dispatch-handler error to the same waiter. - reply_target: ResponseReplyTarget, + reply_target: Box, } impl std::fmt::Debug for ResponseRouter { @@ -4741,9 +5341,15 @@ impl ResponseRouter { /// When [`route_with_result`](Self::route_with_result) is called, the response is sent through the oneshot /// channel to the code that originally sent the request. If that receiver was /// dropped, the response is discarded because there is no local awaiter left. - fn new(id: RequestId, pending_reply: PendingReply, dispatch: ResponseDispatch) -> Self { + fn new( + id: RequestId, + pending_reply: PendingReply, + dispatch: ResponseDispatch, + frame_bytes: Option, + ) -> Self { let PendingReply { method, + metadata_bytes: _, role_id, sender, cancellation_disarm, @@ -4756,6 +5362,7 @@ impl ResponseRouter { sender: Arc::new(Mutex::new(Some(sender))), ordering, dispatch, + frame_bytes, }; let send_target = reply_target.clone(); // A response for the request reached this router, so the request is @@ -4779,7 +5386,7 @@ impl ResponseRouter { send_target.route(response); Ok(()) }), - reply_target, + reply_target: Box::new(reply_target), } } @@ -5474,12 +6081,14 @@ pub struct SentRequest { #[derive(Clone, Debug)] pub(crate) struct SentRequestCancellationDisarm { armed: Arc, + published: Arc, } impl SentRequestCancellationDisarm { fn new() -> Self { Self { armed: Arc::new(AtomicBool::new(true)), + published: Arc::new(AtomicBool::new(false)), } } @@ -5490,6 +6099,8 @@ impl SentRequestCancellationDisarm { struct SentRequestCancellation { message_tx: OutgoingMessageTx, + pending_replies: PendingRepliesRegistrar, + retain_pending_on_drop: AtomicBool, remote_style: crate::role::RemoteStyle, request_id: RequestId, disarm: SentRequestCancellationDisarm, @@ -5498,11 +6109,14 @@ struct SentRequestCancellation { impl SentRequestCancellation { fn new( message_tx: OutgoingMessageTx, + pending_replies: PendingRepliesRegistrar, remote_style: crate::role::RemoteStyle, request_id: RequestId, ) -> Self { Self { message_tx, + pending_replies, + retain_pending_on_drop: AtomicBool::new(false), remote_style, request_id, disarm: SentRequestCancellationDisarm::new(), @@ -5537,6 +6151,11 @@ impl Drop for SentRequestCancellation { if let Err(error) = self.send() { tracing::debug!(?error, "failed to auto-cancel dropped request"); } + // The receiver is gone now; waiting for a peer response would retain + // the pending method and map key without any possible consumer. + if !self.retain_pending_on_drop.load(Ordering::Acquire) { + self.pending_replies.discard_abandoned(&self.request_id); + } } } @@ -5616,7 +6235,7 @@ impl SentRequest { fn new( id: RequestId, method: String, - task_tx: mpsc::UnboundedSender, + task_tx: TaskTx, response_rx: oneshot::Receiver, cancellation: SentRequestCancellation, response_ordering: ResponseOrdering, @@ -5649,6 +6268,11 @@ impl SentRequest { /// handle while automatic cancellation is armed. pub fn detach(self) { self.cancellation.disarm(); + // A detached request must stay registered until it has been + // published: the outgoing actor skips unregistered requests. + self.cancellation + .retain_pending_on_drop + .store(true, Ordering::Release); } /// Send a `$/cancel_request` notification for this outgoing request. @@ -5870,7 +6494,11 @@ impl SentRequest { .await; match response { - Ok(ResponsePayload { result, ack_tx }) => { + Ok(ResponsePayload { + result, + ack_tx, + retained_bytes, + }) => { // Convert the result using to_result for Ok values let typed_result = match result { Ok(json_value) => to_result(json_value), @@ -5878,6 +6506,7 @@ impl SentRequest { }; let outcome = handle(Ok(typed_result)).await; + drop(retained_bytes); // Ack AFTER the handler completes - this is the key // difference from block_task. The dispatch loop waits for @@ -5974,6 +6603,7 @@ impl SentRequest { Ok(ResponsePayload { result: Ok(json_value), ack_tx, + retained_bytes: _, }) => { // Blocking consumers ack before converting or returning the // value, so dispatch can continue while the caller processes it. @@ -5988,6 +6618,7 @@ impl SentRequest { Ok(ResponsePayload { result: Err(err), ack_tx, + retained_bytes: _, }) => { if let Some(tx) = ack_tx { let _ = tx.send(()); @@ -6020,11 +6651,16 @@ impl SentRequest { .await; let (result, ack_tx) = match response { - Ok(ResponsePayload { result, ack_tx }) => { + Ok(ResponsePayload { + result, + ack_tx, + retained_bytes, + }) => { let typed_result = match result { Ok(json_value) => (self.to_result)(json_value), Err(error) => Err(error), }; + drop(retained_bytes); (typed_result, ack_tx) } Err(error) => ( @@ -6339,8 +6975,9 @@ where } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { - self.into_channel_transport() + fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { + let (channel, driver) = self.into_channel_transport(); + (channel, crate::ConnectionDriver::new(driver)) } } @@ -6404,11 +7041,9 @@ where impl futures::Sink + Send + 'static, impl futures::Stream> + Send + 'static, > { - use futures::AsyncBufReadExt; - use futures::io::BufReader; let Self { outgoing, incoming } = self; - let incoming_lines = Box::pin(BufReader::new(incoming).lines()); + let incoming_lines = Box::pin(transport_actor::bounded_lines(Box::pin(incoming))); let outgoing_lines = futures::sink::unfold(Box::pin(outgoing), async move |mut writer, line: String| { write_line(&mut writer, line).await?; @@ -6440,7 +7075,7 @@ where ConnectTo::::connect_to(self.into_lines(), client).await } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { + fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { ConnectTo::::into_channel_and_future(self.into_lines()) } } @@ -6470,122 +7105,1844 @@ where #[derive(Debug)] pub struct Channel { /// Receives frames from the counterpart. - pub rx: mpsc::UnboundedReceiver, + pub rx: FrameReceiver, /// Sends frames to the counterpart. - pub tx: mpsc::UnboundedSender, + pub tx: FrameSender, } -impl Channel { - /// Create a pair of connected channel endpoints. - /// - /// Frames sent through either endpoint are received by the other endpoint. - #[must_use] - pub fn duplex() -> (Self, Self) { - let (a_tx, b_rx) = mpsc::unbounded(); - let (b_tx, a_rx) = mpsc::unbounded(); +/// The byte charge for a frame. Clones refer to the same charge; releasing it +/// requires dropping *every* copy, including deferred dispatch/writer copies. +#[derive(Clone, Debug)] +pub struct FramePermit { + inner: Arc, + additional: Vec, +} + +impl FramePermit { + /// Number of bytes held until every copy of this permit is dropped. + pub fn charged_bytes(&self) -> usize { + self.inner.bytes.load(Ordering::Acquire) + + self + .additional + .iter() + .map(Self::charged_bytes) + .sum::() + } + + /// Reserve separately measured metadata retained after consuming this frame. + /// + /// This reserves `bytes` in every distinct budget covering the source frame, + /// including imported frames. It does not share or release the payload's + /// charge. Retain the returned permit with the metadata, then drop the + /// source permit once its payload has been consumed. + /// + /// Metadata uses data capacity, never the response/cancellation reserve. + /// Failure is immediate and releases any partial reservations; waiting + /// here could deadlock on capacity held by the source frame itself. + pub fn try_reserve_metadata(&self, bytes: usize) -> Result { + fn collect_budgets<'a>(permit: &'a FramePermit, budgets: &mut Vec<&'a Arc>) { + if !budgets + .iter() + .any(|budget| Arc::ptr_eq(budget, &permit.inner.budget)) + { + budgets.push(&permit.inner.budget); + } + for additional in &permit.additional { + collect_budgets(additional, budgets); + } + } - (Self { rx: a_rx, tx: a_tx }, Self { rx: b_rx, tx: b_tx }) + let reserve = |budget: &Arc| { + budget.try_reserve(bytes, true).ok_or_else(|| { + crate::Error::invalid_request().data("retained metadata byte capacity exceeded") + }) + }; + let mut budgets = Vec::new(); + collect_budgets(self, &mut budgets); + let mut budgets = budgets.into_iter(); + let mut permit = reserve(budgets.next().expect("source frame has a budget"))?; + for budget in budgets { + permit.join(reserve(budget)?); + } + Ok(permit) } - /// Copy frames from `rx` to `tx` until the input closes. - /// - /// # Errors - /// - /// Returns an error if the receiving endpoint closes before the input. - pub(crate) async fn copy(mut self) -> Result<(), crate::Error> { - while let Some(frame) = self.rx.next().await { + fn join(&mut self, other: FramePermit) { + self.additional.push(other); + } + + fn charged_for(&self, budget: &Arc) -> usize { + let own = if Arc::ptr_eq(&self.inner.budget, budget) { + self.inner.bytes.load(Ordering::Acquire) + } else { + 0 + }; + own + self + .additional + .iter() + .map(|permit| permit.charged_for(budget)) + .sum::() + } + + fn cover_budget( + &self, + budget: &Arc, + bytes: usize, + data: bool, + ) -> Result<(), crate::Error> { + if Arc::ptr_eq(&self.inner.budget, budget) { + self.cover_frame(bytes, data) + } else { + self.additional + .iter() + .find(|permit| permit.charged_for(budget) > 0) + .expect("destination charge exists") + .cover_budget(budget, bytes, data) + } + } + + fn cover_frame(&self, bytes: usize, data: bool) -> Result<(), crate::Error> { + let budget = &self.inner.budget; + // Aggregated batch permits can exceed the maximum size of one frame. + // That must not allow an oversized frame to bypass the per-frame limit. + if bytes > budget.limits.max_frame_bytes { + return Err( + crate::Error::invalid_request().data("frame exceeds connection byte budget") + ); + } + if data && !self.inner.data { + return Err(crate::Error::invalid_request() + .data("data frame cannot grow a control reservation")); + } + let charged = self.charged_for(budget); + if bytes <= charged { + return Ok(()); + } + let delta = bytes - charged; + let mut state = budget.state.lock().expect("frame budget poisoned"); + if state + .used + .checked_add(delta) + .is_none_or(|used| used > budget.limits.max_queued_bytes) + || data + && state.data_used.checked_add(delta).is_none_or(|used| { + used > budget + .limits + .max_queued_bytes + .saturating_sub(budget.limits.max_frame_bytes) + }) + { + return Err(crate::Error::invalid_request() + .data("outgoing frame exceeds admitted byte capacity")); + } + state.used += delta; + if self.inner.data { + state.data_used += delta; + } + self.inner.bytes.fetch_add(delta, Ordering::Release); + Ok(()) + } +} + +#[derive(Debug)] +struct FramePermitInner { + budget: Arc, + bytes: std::sync::atomic::AtomicUsize, + data: bool, +} + +impl Drop for FramePermitInner { + fn drop(&mut self) { + let mut state = self.budget.state.lock().expect("frame budget poisoned"); + let bytes = self.bytes.load(Ordering::Acquire); + state.used -= bytes; + if self.data { + state.data_used -= bytes; + } + let waiters = state + .waiters + .iter() + .map(|(_, waker)| waker.clone()) + .collect::>(); + drop(state); + for waker in waiters { + waker.wake(); + } + } +} + +#[derive(Debug)] +struct FrameBudget { + limits: ConnectionLimits, + state: Mutex, +} + +#[derive(Debug, Default)] +struct FrameBudgetState { + used: usize, + data_used: usize, + waiters: Vec<(usize, Waker)>, + next_waiter: usize, +} + +struct FrameWaiter { + budget: Arc, + id: Option, +} + +impl Drop for FrameWaiter { + fn drop(&mut self) { + if let Some(id) = self.id { + self.budget + .state + .lock() + .expect("frame budget poisoned") + .waiters + .retain(|(registered, _)| *registered != id); + } + } +} + +impl FrameBudget { + fn try_reserve(self: &Arc, bytes: usize, data: bool) -> Option { + let mut state = self.state.lock().expect("frame budget poisoned"); + if bytes > self.limits.max_frame_bytes + || state.used.checked_add(bytes)? > self.limits.max_queued_bytes + || data + && state.data_used.checked_add(bytes)? + > self + .limits + .max_queued_bytes + .saturating_sub(self.limits.max_frame_bytes) + { + return None; + } + state.used += bytes; + if data { + state.data_used += bytes; + } + Some(FramePermit { + inner: Arc::new(FramePermitInner { + budget: self.clone(), + bytes: std::sync::atomic::AtomicUsize::new(bytes), + data, + }), + additional: Vec::new(), + }) + } + + async fn reserve( + self: &Arc, + bytes: usize, + data: bool, + ) -> Result { + if bytes > self.limits.max_frame_bytes || bytes > self.limits.max_queued_bytes { + return Err( + crate::Error::invalid_request().data("frame exceeds connection byte budget") + ); + } + if data + && bytes + > self + .limits + .max_queued_bytes + .saturating_sub(self.limits.max_frame_bytes) + { + return Err( + crate::Error::invalid_request().data("data frame exceeds connection byte budget") + ); + } + let mut waiter = FrameWaiter { + budget: self.clone(), + id: None, + }; + future::poll_fn(|cx| { + let mut state = self.state.lock().expect("frame budget poisoned"); + if state + .used + .checked_add(bytes) + .is_some_and(|used| used <= self.limits.max_queued_bytes) + && (!data + || state.data_used.checked_add(bytes).is_some_and(|used| { + used <= self + .limits + .max_queued_bytes + .saturating_sub(self.limits.max_frame_bytes) + })) + { + state.used += bytes; + if data { + state.data_used += bytes; + } + Poll::Ready(Ok(FramePermit { + inner: Arc::new(FramePermitInner { + budget: self.clone(), + bytes: std::sync::atomic::AtomicUsize::new(bytes), + data, + }), + additional: Vec::new(), + })) + } else { + if let Some(id) = waiter.id { + let (_, waker) = state + .waiters + .iter_mut() + .find(|(registered, _)| *registered == id) + .expect("registered budget waiter"); + waker.clone_from(cx.waker()); + } else { + let id = state.next_waiter; + state.next_waiter = state.next_waiter.wrapping_add(1); + state.waiters.push((id, cx.waker().clone())); + waiter.id = Some(id); + } + Poll::Pending + } + }) + .await + } +} + +/// A frame with its retained byte admission. Forward this envelope rather +/// than extracting the frame when placing data into another queue. +#[derive(Debug)] +pub struct BudgetedFrame { + frame: TransportFrame, + permit: FramePermit, +} + +impl BudgetedFrame { + /// Borrow the frame without releasing admission. + #[must_use] + pub fn frame(&self) -> &TransportFrame { + &self.frame + } + + /// Borrow the charge when retaining metadata derived from this frame. + /// Cloning the permit retains admission without cloning the payload. + #[must_use] + pub fn permit(&self) -> &FramePermit { + &self.permit + } + + /// Separate the frame and permit for deferred processing. Keep the permit + /// alongside any deferred output until that output has been consumed. + #[must_use] + pub fn into_parts(self) -> (TransportFrame, FramePermit) { + (self.frame, self.permit) + } + + /// Release the frame's admission explicitly after consuming it. + #[must_use] + pub fn into_frame(self) -> TransportFrame { + self.frame + } +} + +/// Pollable receive half of an in-memory duplex. +#[derive(Debug)] +pub struct FrameReceiver { + rx: std::pin::Pin>>, + _slots: async_channel::Sender<()>, +} + +#[cfg(feature = "unstable_mcp_over_acp")] +impl FrameReceiver { + /// Reject new output while retaining already-accepted frames for draining. + pub(crate) fn close(&self) { + self.rx.close(); + } +} + +/// A slot is reserved before a frame enters the queue and returned at dequeue. +/// Its drop also returns reservations abandoned by a cancelled send or sink. +#[derive(Debug)] +struct FrameSlot(async_channel::Sender<()>); + +impl Drop for FrameSlot { + fn drop(&mut self) { + let _ = self.0.try_send(()); + } +} + +#[derive(Debug)] +struct QueuedFrame { + frame: BudgetedFrame, + slot: FrameSlot, +} + +impl futures::Stream for FrameReceiver { + type Item = BudgetedFrame; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut Context<'_>, + ) -> Poll> { + self.rx.as_mut().poll_next(cx).map(|item| { + item.map(|queued| { + drop(queued.slot); + queued.frame + }) + }) + } +} + +/// Backpressured frame sink. A synchronous send fails when its finite queue is full; +/// asynchronous producers should use [`SinkExt::send`] instead. +pub struct FrameSender { + tx: async_channel::Sender, + slots: Box>, + slot_return: async_channel::Sender<()>, + budget: Arc, + ready: Option, + waiting: Mutex>>>, +} + +impl std::fmt::Debug for FrameSender { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("FrameSender") + .field("tx", &self.tx) + .field("budget", &self.budget) + .finish_non_exhaustive() + } +} + +/// Shared byte admission independent of a channel's send half. Used when an +/// adapter stages frames before forwarding them to the channel sink. +#[derive(Clone, Debug)] +pub struct FrameAdmission(Arc); + +impl FrameAdmission { + /// Limits shared by both halves of the duplex connection. + #[must_use] + pub fn limits(&self) -> ConnectionLimits { + self.0.limits + } + + fn try_reserve_bytes(&self, bytes: usize, data: bool) -> Option { + self.0.try_reserve(bytes, data) + } + + async fn reserve_bytes(&self, bytes: usize, data: bool) -> Result { + self.0.reserve(bytes, data).await + } + + /// Admit a frame before placing it into any staging queue. + pub fn try_admit(&self, frame: TransportFrame) -> Result { + let bytes = frame + .to_json() + .map_err(|_| FrameSendError { + frame: Box::new(frame.clone()), + reason: "cannot serialize outgoing JSON-RPC frame", + })? + .len(); + let Some(permit) = self.0.try_reserve(bytes, !frame.is_control()) else { + return Err(FrameSendError { + frame: Box::new(frame), + reason: "outgoing frame byte capacity exceeded", + }); + }; + Ok(BudgetedFrame { frame, permit }) + } + + /// Wait for byte capacity when staging a frame outside inline dispatch. + pub async fn admit(&self, frame: TransportFrame) -> Result { + let bytes = frame.to_json()?.len(); + let permit = self.0.reserve(bytes, !frame.is_control()).await?; + Ok(BudgetedFrame { frame, permit }) + } +} + +impl Clone for FrameSender { + fn clone(&self) -> Self { + Self { + tx: self.tx.clone(), + slots: Box::new((*self.slots).clone()), + slot_return: self.slot_return.clone(), + budget: self.budget.clone(), + ready: None, + waiting: Mutex::new(None), + } + } +} + +/// Failure to admit a frame (capacity, size, or closed receiver). The original +/// frame remains available to callers; nothing is silently discarded. +#[derive(Debug)] +pub struct FrameSendError { + frame: Box, + reason: &'static str, +} + +impl std::fmt::Display for FrameSendError { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.write_str(self.reason) + } +} + +impl std::error::Error for FrameSendError {} + +impl FrameSendError { + /// Recover the frame that was not admitted. + pub fn into_inner(self) -> TransportFrame { + *self.frame + } +} + +impl FrameSender { + async fn wait_for_slot( + tx: async_channel::Sender, + slots: async_channel::Receiver<()>, + slot_return: async_channel::Sender<()>, + ) -> Result { + if tx.is_closed() { + return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); + } + match future::select(Box::pin(slots.recv()), Box::pin(tx.closed())).await { + Either::Left((Ok(()), _)) if !tx.is_closed() => Ok(FrameSlot(slot_return)), + _ => Err(crate::Error::invalid_request().data("outgoing frame queue closed")), + } + } + + fn try_slot(&self) -> Result { + if self.tx.is_closed() { + return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); + } + self.slots + .try_recv() + .map(|()| FrameSlot(self.slot_return.clone())) + .map_err(|_| { + crate::Error::invalid_request().data("outgoing frame queue full or closed") + }) + } + + async fn reserve_slot(&self) -> Result { + Self::wait_for_slot( + self.tx.clone(), + (*self.slots).clone(), + self.slot_return.clone(), + ) + .await + } + + fn enqueue(&self, frame: BudgetedFrame, slot: FrameSlot) -> Result<(), crate::Error> { + self.tx + .try_send(QueuedFrame { frame, slot }) + .map_err(crate::util::internal_error) + } + + /// Obtain the byte admission handle without retaining this channel sender. + pub fn admission(&self) -> FrameAdmission { + FrameAdmission(self.budget.clone()) + } + + /// Fail immediately rather than blocking a protocol dispatcher on its own output. + pub fn try_send(&self, frame: TransportFrame) -> Result<(), FrameSendError> { + let Ok(slot) = self.try_slot() else { + return Err(FrameSendError { + frame: Box::new(frame), + reason: "outgoing frame queue full or closed", + }); + }; + let budgeted = self.admission().try_admit(frame)?; + self.tx + .try_send(QueuedFrame { + frame: budgeted, + slot, + }) + .map_err(|error| FrameSendError { + frame: Box::new(error.into_inner().frame.frame), + reason: "outgoing frame queue full or closed", + }) + } + + /// Await byte and frame capacity outside ordered dispatch. + pub async fn send_frame(&self, frame: TransportFrame) -> Result<(), crate::Error> { + let bytes = frame.to_json()?.len(); + let permit = match future::select( + Box::pin(self.budget.reserve(bytes, !frame.is_control())), + Box::pin(self.tx.closed()), + ) + .await + { + Either::Left((result, _)) => result?, + Either::Right(((), _)) => { + return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); + } + }; + let slot = self.reserve_slot().await?; + self.enqueue(BudgetedFrame { frame, permit }, slot) + } + + /// Transfer an application message's charge into its framed representation. + /// Any transform expansion must grow that lease immediately, rather than + /// awaiting capacity held by this very message. + async fn send_admitted( + &self, + frame: TransportFrame, + permit: FramePermit, + ) -> Result<(), crate::Error> { + let bytes = frame.to_json()?.len(); + permit.cover_frame(bytes, !frame.is_control())?; + let slot = self.reserve_slot().await?; + self.enqueue(BudgetedFrame { frame, permit }, slot) + } + + /// Stop accepting frames on this queue. + pub fn close_channel(&self) { + self.tx.close(); + } + + /// Return whether the receiving endpoint has closed. + pub fn is_closed(&self) -> bool { + self.tx.is_closed() + } +} + +impl Sink for FrameSender { + type Error = crate::Error; + + fn poll_ready( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + let this = self.get_mut(); + if this.tx.is_closed() { + this.ready.take(); + this.waiting + .get_mut() + .expect("frame sender poisoned") + .take(); + return Poll::Ready(Err( + crate::Error::invalid_request().data("outgoing frame queue closed") + )); + } + if this.ready.is_some() { + return Poll::Ready(Ok(())); + } + let waiting = this.waiting.get_mut().expect("frame sender poisoned"); + if waiting.is_none() { + *waiting = Some(Box::pin(Self::wait_for_slot( + this.tx.clone(), + (*this.slots).clone(), + this.slot_return.clone(), + ))); + } + match waiting.as_mut().expect("slot waiter").as_mut().poll(cx) { + Poll::Pending => Poll::Pending, + Poll::Ready(result) => { + waiting.take(); + match result { + Ok(slot) => { + this.ready = Some(slot); + Poll::Ready(Ok(())) + } + Err(error) => Poll::Ready(Err(error)), + } + } + } + } + + fn start_send( + self: std::pin::Pin<&mut Self>, + mut item: BudgetedFrame, + ) -> Result<(), Self::Error> { + let this = self.get_mut(); + let slot = this + .ready + .take() + .ok_or_else(|| crate::Error::invalid_request().data("frame sender not ready"))?; + if this.tx.is_closed() { + return Err(crate::Error::invalid_request().data("outgoing frame queue closed")); + } + let bytes = item.frame.to_json()?.len(); + let data = !item.frame.is_control(); + if bytes > this.budget.limits.max_frame_bytes { + return Err( + crate::Error::invalid_request().data("frame exceeds connection byte budget") + ); + } + if item.permit.charged_for(&this.budget) > 0 { + item.permit.cover_budget(&this.budget, bytes, data)?; + } else { + let permit = this.budget.try_reserve(bytes, data).ok_or_else(|| { + crate::Error::invalid_request().data("outgoing frame byte capacity exceeded") + })?; + item.permit.join(permit); + } + this.enqueue(item, slot) + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + _cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + let this = self.get_mut(); + Poll::Ready(if this.tx.is_closed() { + Err(crate::Error::invalid_request().data("outgoing frame queue closed")) + } else { + Ok(()) + }) + } + + fn poll_close( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + let this = self.get_mut(); + match std::pin::Pin::new(&mut *this).poll_flush(cx) { + Poll::Ready(Ok(())) => { + this.tx.close(); + Poll::Ready(Ok(())) + } + other => other, + } + } +} + +impl Channel { + /// Create a pair of connected channel endpoints. + /// + /// Frames sent through either endpoint are received by the other endpoint. + #[must_use] + pub fn duplex() -> (Self, Self) { + Self::duplex_with_limits(ConnectionLimits::default()) + } + + /// Create a connected pair sharing one finite byte budget. + #[must_use] + pub fn duplex_with_limits(limits: ConnectionLimits) -> (Self, Self) { + let budget = Arc::new(FrameBudget { + limits, + state: Mutex::new(FrameBudgetState::default()), + }); + let (a_tx, b_rx) = async_channel::bounded(limits.max_queued_frames.max(1)); + let (b_tx, a_rx) = async_channel::bounded(limits.max_queued_frames.max(1)); + let (a_slot_tx, a_slots) = async_channel::bounded(limits.max_queued_frames.max(1)); + let (b_slot_tx, b_slots) = async_channel::bounded(limits.max_queued_frames.max(1)); + for _ in 0..limits.max_queued_frames.max(1) { + a_slot_tx.try_send(()).expect("initial queue slot"); + b_slot_tx.try_send(()).expect("initial queue slot"); + } + ( + Self { + rx: FrameReceiver { + rx: Box::pin(a_rx), + _slots: b_slot_tx.clone(), + }, + tx: FrameSender { + tx: a_tx, + slots: Box::new(a_slots), + slot_return: a_slot_tx.clone(), + budget: budget.clone(), + ready: None, + waiting: Mutex::new(None), + }, + }, + Self { + rx: FrameReceiver { + rx: Box::pin(b_rx), + _slots: a_slot_tx.clone(), + }, + tx: FrameSender { + tx: b_tx, + slots: Box::new(b_slots), + slot_return: b_slot_tx, + budget, + ready: None, + waiting: Mutex::new(None), + }, + }, + ) + } + + /// Copy frames from `rx` to `tx` until the input closes. + /// + /// # Errors + /// + /// Returns an error if the receiving endpoint closes before the input. + pub(crate) async fn copy(mut self) -> Result<(), crate::Error> { + while let Some(frame) = self.rx.next().await { self.tx - .unbounded_send(frame) + .send(frame) + .await .map_err(crate::util::internal_error)?; } Ok(()) } - /// Bridge two endpoints while inspecting every valid message. - /// - /// Observers are invoked in source order, including for each valid member of - /// a batch. The original frame is forwarded unchanged after inspection. - /// - /// # Errors - /// - /// Returns an observer error or an error if a destination closes before its - /// source. - pub async fn bridge_with_inspection( - left: Self, - right: Self, - mut left_to_right: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send, - mut right_to_left: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send, - ) -> Result<(), crate::Error> { - let Self { - rx: mut left_rx, - tx: left_tx, - } = left; - let Self { - rx: mut right_rx, - tx: right_tx, - } = right; + /// Bridge two endpoints while inspecting every valid message. + /// + /// Observers are invoked in source order, including for each valid member of + /// a batch. The original frame is forwarded unchanged after inspection. + /// + /// # Errors + /// + /// Returns an observer error or an error if a destination closes before its + /// source. + pub async fn bridge_with_inspection( + left: Self, + right: Self, + mut left_to_right: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send, + mut right_to_left: impl FnMut(&RawJsonRpcMessage) -> Result<(), crate::Error> + Send, + ) -> Result<(), crate::Error> { + let Self { + rx: mut left_rx, + tx: mut left_tx, + } = left; + let Self { + rx: mut right_rx, + tx: mut right_tx, + } = right; + + let left_to_right = async move { + while let Some(frame) = left_rx.next().await { + frame.frame().inspect_messages(&mut left_to_right)?; + right_tx + .send(frame) + .await + .map_err(crate::util::internal_error)?; + } + Ok::<(), crate::Error>(()) + }; + let right_to_left = async move { + while let Some(frame) = right_rx.next().await { + frame.frame().inspect_messages(&mut right_to_left)?; + left_tx + .send(frame) + .await + .map_err(crate::util::internal_error)?; + } + Ok::<(), crate::Error>(()) + }; + + futures::try_join!(left_to_right, right_to_left)?; + Ok(()) + } +} + +impl ConnectTo for Channel { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), crate::Error> { + let (client_channel, client_future) = client.into_channel_and_future(); + let passive = client_future.is_passive(); + + let outbound = Channel { + rx: client_channel.rx, + tx: self.tx, + } + .copy(); + let inbound = Channel { + rx: self.rx, + tx: client_channel.tx, + } + .copy(); + if passive { + // Neither channel owns the remote application. Preserve half-close: + // input EOF must still allow responses to drain the other way. + futures::try_join!(inbound, outbound)?; + return Ok(()); + } + // Poll output while the client is running: its requests may be needed + // to let either peer finish. A raw Channel has a no-op driver, so driver + // completion alone is not a signal to stop forwarding its input. + let local = async move { + futures::try_join!(client_future, outbound)?; + Ok::<(), crate::Error>(()) + }; + match future::select(Box::pin(local), Box::pin(inbound)).await { + Either::Left((result, _inbound)) => { + // The local client has finished and its accepted output drained. + // Do not also wait for a remote sender that can remain alive. + result + } + Either::Right((result, local)) => { + result?; + local.await + } + } + } + + fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { + (self, crate::ConnectionDriver::passive()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + struct SendOneThenFinish; + + impl ConnectTo for SendOneThenFinish { + async fn connect_to( + self, + peer: impl ConnectTo, + ) -> Result<(), crate::Error> { + let (channel, driver) = peer.into_channel_and_future(); + channel + .tx + .send_frame(TransportFrame::Single(RawJsonRpcMessage::notification( + "finished".into(), + serde_json::json!({}), + )?)) + .await + .map_err(crate::util::internal_error)?; + drop(channel); + driver.await + } + } + + #[tokio::test] + async fn channel_connect_finishes_without_remote_eof_after_local_drain() { + let (local, mut remote) = Channel::duplex(); + let connection = tokio::spawn(ConnectTo::::connect_to( + local, + SendOneThenFinish, + )); + let frame = tokio::time::timeout(std::time::Duration::from_secs(2), remote.rx.next()) + .await + .expect("accepted frame should arrive") + .expect("channel open"); + assert!(matches!( + frame.frame(), + TransportFrame::Single(RawJsonRpcMessage::Notification(_)) + )); + tokio::time::timeout(std::time::Duration::from_secs(2), connection) + .await + .expect("local completion must not wait for remote sender") + .expect("connection task") + .expect("connection result"); + // Retaining the remote sender did not block local completion. The + // completed endpoint has now closed its receiving half. + assert!(remote.tx.is_closed()); + } + + #[tokio::test] + async fn frame_permits_survive_dequeue_until_consumed() { + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("capacity".into(), serde_json::json!({})).unwrap(), + ); + let bytes = frame.to_json().unwrap().len(); + let (left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes * 2, + max_queued_bytes: bytes * 4, + max_queued_frames: 8, + }); + left.tx.try_send(frame.clone()).unwrap(); + left.tx.try_send(frame.clone()).unwrap(); + let held = right.rx.next().await.unwrap(); + assert_eq!( + held.frame().to_json().unwrap().len(), + held.permit.charged_bytes() + ); + assert!( + left.tx.try_send(frame.clone()).is_err(), + "dequeue must not release byte admission" + ); + drop(held); + left.tx + .try_send(frame) + .expect("dropping the last permit releases capacity"); + } + + #[cfg(feature = "unstable_mcp_over_acp")] + #[test] + fn receiver_close_drains_accepted_output_and_wakes_blocked_senders() { + let (source, mut destination) = Channel::duplex_with_limits(ConnectionLimits { + max_queued_frames: 1, + ..ConnectionLimits::default() + }); + let frame = capacity_frame(); + let escaped = source.tx.clone(); + source.tx.try_send(frame.clone()).unwrap(); + let mut blocked = Box::pin(escaped.send_frame(frame.clone())); + assert!(blocked.as_mut().now_or_never().is_none()); + + destination.rx.close(); + assert!( + blocked.now_or_never().unwrap().is_err(), + "receiver closure must wake a blocked producer" + ); + assert!(escaped.try_send(frame.clone()).is_err()); + let accepted = destination.rx.next().now_or_never().unwrap().unwrap(); + assert_eq!( + accepted.frame().to_json().unwrap(), + frame.to_json().unwrap() + ); + assert!( + destination.rx.next().now_or_never().unwrap().is_none(), + "draining must not wait for the escaped sender to be dropped" + ); + } + + fn capacity_frame() -> TransportFrame { + TransportFrame::Single( + RawJsonRpcMessage::notification("capacity".into(), serde_json::json!({})).unwrap(), + ) + } + + #[test] + fn cloned_frame_senders_do_not_expand_queue_capacity() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes * 2, + max_queued_bytes: bytes * 5000, + max_queued_frames: 2, + }); + let clones = (0..3000).map(|_| left.tx.clone()).collect::>(); + clones[0].try_send(frame.clone()).unwrap(); + clones[1].try_send(frame.clone()).unwrap(); + for tx in &clones { + assert!(tx.try_send(frame.clone()).is_err()); + } + drop(right.rx.next().now_or_never().unwrap()); + clones[2999] + .try_send(frame) + .expect("one dequeue restores precisely one slot"); + } + + #[test] + fn task_and_dynamic_queues_remain_bounded_across_clones() { + for name in ["task", "dynamic"] { + let (tx, mut rx) = admission::channel_with_capacity::(2); + let clones = (0..3000).map(|_| tx.clone()).collect::>(); + clones[0].unbounded_send(0).unwrap(); + clones[1].unbounded_send(1).unwrap(); + assert!( + clones + .iter() + .all(|sender| sender.unbounded_send(2).is_err()), + "{name}" + ); + assert_eq!(rx.next().now_or_never().unwrap(), Some(0)); + clones[2999] + .unbounded_send(3) + .expect("dequeue restores one slot"); + } + } + + #[tokio::test] + async fn imported_frames_obey_destination_frame_and_byte_limits() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (source, mut source_peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes * 2, + max_queued_bytes: bytes * 8, + max_queued_frames: 2, + }); + source.tx.try_send(frame.clone()).unwrap(); + let held = source_peer.rx.next().await.unwrap(); + let lease = held.permit.clone(); + let source_budget = source.tx.budget.clone(); + let (mut smaller, _peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes - 1, + max_queued_bytes: bytes * 8, + max_queued_frames: 2, + }); + assert!(smaller.tx.send(held).await.is_err()); + assert_eq!(source_budget.state.lock().unwrap().used, bytes); + drop(lease); + assert_eq!(source_budget.state.lock().unwrap().used, 0); + + source.tx.try_send(frame.clone()).unwrap(); + let held = source_peer.rx.next().await.unwrap(); + let lease = held.permit.clone(); + let (mut smaller, _peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2 - 1, + max_queued_frames: 2, + }); + assert!(smaller.tx.send(held).await.is_err()); + assert_eq!(source_budget.state.lock().unwrap().used, bytes); + drop(lease); + assert_eq!(source_budget.state.lock().unwrap().used, 0); + source + .tx + .try_send(frame) + .expect("rejected import releases source lease"); + } + + #[tokio::test] + async fn same_budget_frame_handoff_does_not_charge_twice() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (mut left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2, + max_queued_frames: 2, + }); + left.tx.try_send(frame).unwrap(); + let held = right.rx.next().await.unwrap(); + right.tx.send(held).await.unwrap(); + let held = left.rx.next().await.unwrap(); + assert_eq!(left.tx.budget.state.lock().unwrap().used, bytes); + drop(held); + assert_eq!(left.tx.budget.state.lock().unwrap().used, 0); + } + + #[tokio::test] + async fn imported_frame_retains_independent_budget_charges_without_recharging_on_return() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let limits = ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2, + max_queued_frames: 2, + }; + let (source, mut source_peer) = Channel::duplex_with_limits(limits); + let (mut destination, mut destination_peer) = Channel::duplex_with_limits(limits); + source.tx.try_send(frame).unwrap(); + destination + .tx + .send(source_peer.rx.next().await.unwrap()) + .await + .unwrap(); + assert_eq!(source.tx.budget.state.lock().unwrap().used, bytes); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, bytes); + let imported = destination_peer.rx.next().await.unwrap(); + destination_peer.tx.send(imported).await.unwrap(); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, bytes); + drop(destination.rx.next().await.unwrap()); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 0); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 0); + } + + #[test] + fn cancelled_byte_waiters_are_unregistered() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, _right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2, + max_queued_frames: 2, + }); + let _held = left.tx.admission().try_admit(frame.clone()).unwrap(); + for _ in 0..1000 { + { + let waiting = left.tx.budget.reserve(bytes, true); + futures::pin_mut!(waiting); + assert!(waiting.as_mut().now_or_never().is_none()); + } + assert!(left.tx.budget.state.lock().unwrap().waiters.is_empty()); + } + } + + #[tokio::test] + async fn closing_receiver_wakes_byte_blocked_sender() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 2, + max_queued_frames: 2, + }); + let _held = left.tx.admission().try_admit(frame.clone()).unwrap(); + let mut waiting = Box::pin(left.tx.send_frame(frame)); + assert!(waiting.as_mut().now_or_never().is_none()); + drop(right); + assert!(waiting.await.is_err()); + assert!(left.tx.budget.state.lock().unwrap().waiters.is_empty()); + } + + #[test] + fn cancelling_queued_async_send_releases_its_byte_charge() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 4, + max_queued_frames: 1, + }); + left.tx.try_send(frame.clone()).unwrap(); + { + let waiting = left.tx.send_frame(frame.clone()); + futures::pin_mut!(waiting); + assert!(waiting.as_mut().now_or_never().is_none()); + assert_eq!(left.tx.budget.state.lock().unwrap().used, bytes * 2); + } + assert_eq!(left.tx.budget.state.lock().unwrap().used, bytes); + drop(right.rx.next().now_or_never().unwrap()); + assert_eq!(left.tx.budget.state.lock().unwrap().used, 0); + } + + #[test] + fn sink_clones_share_slots_and_dequeue_restores_capacity() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, mut right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 100, + max_queued_frames: 1, + }); + let mut clones = (0..32).map(|_| left.tx.clone()).collect::>(); + clones[0] + .feed(left.tx.admission().try_admit(frame.clone()).unwrap()) + .now_or_never() + .expect("first feed must complete") + .unwrap(); + for tx in &mut clones[1..] { + let admitted = left.tx.admission().try_admit(frame.clone()).unwrap(); + assert!(tx.feed(admitted).now_or_never().is_none()); + } + assert!(left.tx.try_send(frame.clone()).is_err()); + drop(right.rx.next().now_or_never().unwrap()); + clones[31] + .feed(left.tx.admission().try_admit(frame).unwrap()) + .now_or_never() + .expect("dequeue restores a slot") + .unwrap(); + } + + #[test] + fn dropping_ready_sink_releases_reserved_slot() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, _right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 8, + max_queued_frames: 1, + }); + let mut reserved = left.tx.clone(); + assert!( + std::pin::Pin::new(&mut reserved) + .poll_ready(&mut Context::from_waker(futures::task::noop_waker_ref())) + .is_ready() + ); + assert!(left.tx.try_send(frame.clone()).is_err()); + drop(reserved); + left.tx + .try_send(frame) + .expect("dropping readiness frees slot"); + } + + #[tokio::test] + async fn closing_receiver_wakes_slot_blocked_sink_and_sender() { + let frame = capacity_frame(); + let bytes = frame.to_json().unwrap().len(); + let (left, right) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: bytes, + max_queued_bytes: bytes * 8, + max_queued_frames: 1, + }); + left.tx.try_send(frame.clone()).unwrap(); + let mut waiting_sink = left.tx.clone(); + let mut waiting_feed = + Box::pin(waiting_sink.feed(left.tx.admission().try_admit(frame.clone()).unwrap())); + assert!(waiting_feed.as_mut().now_or_never().is_none()); + let mut waiting_send = Box::pin(left.tx.send_frame(frame)); + assert!(waiting_send.as_mut().now_or_never().is_none()); + drop(right); + assert!(waiting_feed.await.is_err()); + assert!(waiting_send.await.is_err()); + } + + #[test] + fn retained_metadata_charges_each_source_budget_without_pinning_payload() { + let (source, _source_peer) = Channel::duplex(); + let (destination, _destination_peer) = Channel::duplex(); + let mut payload = source.tx.budget.try_reserve(512, true).unwrap(); + payload.join(payload.clone()); + payload.join(destination.tx.budget.try_reserve(512, true).unwrap()); + + let metadata = payload.try_reserve_metadata(8).unwrap(); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 520); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 520); + drop(payload); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 8); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 8); + + let copy = metadata.clone(); + drop(metadata); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 8); + drop(copy); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 0); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 0); + } + + #[test] + fn retained_metadata_failure_rolls_back_and_cannot_use_control_capacity() { + let (source, _source_peer) = Channel::duplex(); + let (destination, _destination_peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 64, + max_queued_bytes: 128, + max_queued_frames: 4, + }); + let mut payload = source.tx.budget.try_reserve(60, true).unwrap(); + payload.join(destination.tx.budget.try_reserve(60, true).unwrap()); + assert!(payload.try_reserve_metadata(8).is_err()); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 60); + assert_eq!(destination.tx.budget.state.lock().unwrap().used, 60); + + let metadata = payload.try_reserve_metadata(4).unwrap(); + assert!(payload.try_reserve_metadata(1).is_err()); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 64); + drop(metadata); + assert_eq!(source.tx.budget.state.lock().unwrap().used, 60); + } + + fn application_channel( + admission: FrameAdmission, + ) -> ( + outgoing_actor::OutgoingMessageTx, + admission::Receiver, + ) { + admission::budgeted_channel( + admission, + OutgoingMessage::charged_bytes, + OutgoingMessage::with_permit, + OutgoingMessage::is_control, + OutgoingMessage::is_urgent, + ) + } + + #[test] + fn application_byte_wait_is_interrupted_by_receiver_close() { + let message = || OutgoingMessage::Notification { + untyped: UntypedMessage::new("held", serde_json::json!({})).unwrap(), + }; + let charge = message().charged_bytes().unwrap(); + let (channel, _peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: charge, + max_queued_bytes: charge * 2, + max_queued_frames: 2, + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + tx.unbounded_send(message()).unwrap(); + let held = rx.next().now_or_never().unwrap().unwrap(); + let mut blocked = Box::pin(tx.send(message())); + assert!(blocked.as_mut().now_or_never().is_none()); + drop(rx); + assert!( + blocked + .now_or_never() + .expect("closed queue must wake a byte waiter") + .is_err() + ); + // The held payload deliberately outlives closure of its queue. + drop(held); + } + + #[test] + fn unbudgeted_control_queue_reports_its_configured_capacity() { + let (tx, _rx) = admission::channel_with_capacity::(2); + assert_eq!(tx.queue_capacity(), 2); + assert_eq!(tx.clone().queue_capacity(), 2); + } + + #[test] + fn routed_results_have_independent_retained_charges() { + let limits = ConnectionLimits { + max_frame_bytes: 1024, + max_queued_bytes: 2700, + max_queued_frames: 8, + }; + let (channel, _) = Channel::duplex_with_limits(limits); + let admission = channel.tx.admission(); + let initial = admission.try_reserve_bytes(100, true).unwrap(); + let value = serde_json::json!("x".repeat(700)); + let mut receivers = Vec::new(); + for i in 0..2 { + let (sender, receiver) = oneshot::channel(); + let id = RequestId::Str(format!("response-{i}")); + let pending = PendingReply { + method: "test".into(), + metadata_bytes: None, + role_id: crate::role::UntypedRole.role_id(), + sender, + cancellation_disarm: SentRequestCancellationDisarm::new(), + ordering: ResponseOrdering::default(), + response_route_hook: None, + }; + let (dispatch, _) = incoming_actor::dispatch_from_response( + id, + pending, + Ok(value.clone()), + Some(initial.clone()), + ); + let Dispatch::Response(result, router) = dispatch else { + panic!("response expected") + }; + router.route_with_result(result).unwrap(); + receivers.push(receiver); + } + drop(initial); + let used = admission.0.state.lock().unwrap().used; + assert_eq!(used, 2 * serde_json::to_vec(&value).unwrap().len()); + let first = futures::executor::block_on(receivers.remove(0)).unwrap(); + assert!(first.result.is_ok()); + assert!(admission.0.state.lock().unwrap().used > 0); + drop(first); + drop(receivers); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + + #[test] + fn oversized_transformed_result_fails_without_waiting_on_its_frame() { + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 1000, + max_queued_frames: 8, + }); + let admission = channel.tx.admission(); + let frame = admission.try_reserve_bytes(200, true).unwrap(); + let (sender, receiver) = oneshot::channel(); + let pending = PendingReply { + method: "transform".into(), + metadata_bytes: None, + role_id: crate::role::UntypedRole.role_id(), + sender, + cancellation_disarm: SentRequestCancellationDisarm::new(), + ordering: ResponseOrdering::default(), + response_route_hook: None, + }; + let (dispatch, _) = incoming_actor::dispatch_from_response( + RequestId::Str("transform".into()), + pending, + Ok(serde_json::json!(null)), + Some(frame.clone()), + ); + let Dispatch::Response(_, router) = dispatch else { + panic!("response expected") + }; + router.route(serde_json::json!("x".repeat(400))).unwrap(); + drop(frame); + let received = futures::executor::block_on(receiver).unwrap(); + assert!(received.result.is_err()); + drop(received); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + + #[test] + fn callback_keeps_result_admitted_until_callback_finishes() { + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 1024, + max_queued_bytes: 2700, + max_queued_frames: 8, + }); + let admission = channel.tx.admission(); + let (message_tx, _message_rx) = application_channel(admission.clone()); + let (task_tx, mut task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); + let pending = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx, + task_tx, + dynamic_handler_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending.registrar(), + ProtocolMode::disabled(), + ); + let sent = connection.send_request_to( + crate::role::UntypedRole, + UntypedMessage::new("callback", serde_json::json!({})).unwrap(), + ); + let id = sent.id().clone(); + let frame = admission.try_reserve_bytes(100, true).unwrap(); + let pending_reply = pending.remove(&id).unwrap(); + let (dispatch, _) = incoming_actor::dispatch_from_response( + id, + pending_reply, + Ok(serde_json::json!("x".repeat(500))), + Some(frame.clone()), + ); + let Dispatch::Response(result, router) = dispatch else { + panic!("response expected") + }; + router.route_with_result(result).unwrap(); + drop(frame); + let (finish_tx, finish_rx) = oneshot::channel::<()>(); + sent.on_receiving_result(move |result| async move { + assert!(result.is_ok()); + finish_rx.await.unwrap(); + Ok(()) + }) + .unwrap(); + let task = futures::FutureExt::now_or_never(futures::StreamExt::next(&mut task_rx)) + .unwrap() + .unwrap(); + let mut running = Box::pin(task.run_for_test()); + assert!(running.as_mut().now_or_never().is_none()); + assert!(admission.0.state.lock().unwrap().used >= 500); + finish_tx.send(()).unwrap(); + futures::executor::block_on(running).unwrap(); + // The outgoing frame is still queued; only the callback's result + // charge has been released. + assert!(admission.0.state.lock().unwrap().used < 500); + } - let left_to_right = async move { - while let Some(frame) = left_rx.next().await { - frame.inspect_messages(&mut left_to_right)?; - right_tx - .unbounded_send(frame) - .map_err(crate::util::internal_error)?; - } - Ok::<(), crate::Error>(()) + #[test] + fn cloned_application_senders_respect_item_capacity_independently_of_bytes() { + let message = || OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({})).unwrap(), }; - let right_to_left = async move { - while let Some(frame) = right_rx.next().await { - frame.inspect_messages(&mut right_to_left)?; - left_tx - .unbounded_send(frame) - .map_err(crate::util::internal_error)?; - } - Ok::<(), crate::Error>(()) + let charge = message().charged_bytes().unwrap(); + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: charge * 2, + max_queued_bytes: charge * 6000, + max_queued_frames: 2, + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + let clones = (0..3000).map(|_| tx.clone()).collect::>(); + clones[0].unbounded_send(message()).unwrap(); + clones[1].unbounded_send(message()).unwrap(); + assert!( + clones + .iter() + .all(|sender| sender.unbounded_send(message()).is_err()) + ); + drop(rx.next().now_or_never().unwrap()); + clones[2999].unbounded_send(message()).unwrap(); + } + + #[test] + fn application_payload_is_charged_after_dequeue_until_dropped() { + let message = OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({"value": "123"})).unwrap(), }; + let charge = message.charged_bytes().unwrap(); + let frame_bytes = charge + 32; + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: frame_bytes, + max_queued_bytes: frame_bytes + 2 * charge, + max_queued_frames: 3, + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({"value": "123"})).unwrap(), + }) + .unwrap(); + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({"value": "123"})).unwrap(), + }) + .unwrap(); + let held = rx.next().now_or_never().unwrap().unwrap(); + assert!( + tx.unbounded_send(message).is_err(), + "dequeue must retain application admission" + ); + drop(held); + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("capacity", serde_json::json!({"value": "123"})).unwrap(), + }) + .expect("capacity is recovered after the retained application message is dropped"); + } - futures::try_join!(left_to_right, right_to_left)?; - Ok(()) + #[tokio::test] + async fn application_lease_moves_into_writer_frame_without_recharging() { + let message = OutgoingMessage::Notification { + untyped: UntypedMessage::new("handoff", serde_json::json!({"value": "abc"})).unwrap(), + }; + let charge = message.charged_bytes().unwrap(); + let frame_bytes = charge + 16; + let (sender, mut receiver) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: frame_bytes, + max_queued_bytes: frame_bytes + charge, + max_queued_frames: 2, + }); + let (tx, mut application_rx) = application_channel(sender.tx.admission()); + tx.unbounded_send(message).unwrap(); + let OutgoingMessage::Admitted { message, permit } = application_rx.next().await.unwrap() + else { + panic!("application admission must wrap the queued payload"); + }; + let OutgoingMessage::Notification { untyped } = *message else { + panic!("expected notification"); + }; + let frame = TransportFrame::Single(untyped.into_raw_jsonrpc_message(None).unwrap()); + sender.tx.send_admitted(frame, permit).await.unwrap(); + let held = receiver.rx.next().await.unwrap(); + assert!( + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("handoff", serde_json::json!({"value": "abc"})) + .unwrap(), + }) + .is_err(), + "writer-held frame keeps its application charge" + ); + drop(held); + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("handoff", serde_json::json!({"value": "abc"})).unwrap(), + }) + .expect("capacity is recovered after the writer frame is released"); } -} -impl ConnectTo for Channel { - async fn connect_to(self, client: impl ConnectTo) -> Result<(), crate::Error> { - let (client_channel, client_future) = client.into_channel_and_future(); + #[test] + fn cancellation_lane_is_ready_when_data_queue_is_full() { + let ordinary = OutgoingMessage::Notification { + untyped: UntypedMessage::new("ordinary", serde_json::json!({})).unwrap(), + }; + let cancel = OutgoingMessage::Notification { + untyped: UntypedMessage::new( + "$/cancel_request", + serde_json::json!({"requestId":"one"}), + ) + .unwrap(), + }; + let frame_bytes = ordinary + .charged_bytes() + .unwrap() + .max(cancel.charged_bytes().unwrap()) + + 16; + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: frame_bytes, + max_queued_bytes: frame_bytes * 2, + max_queued_frames: 1, + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + tx.unbounded_send(ordinary).unwrap(); + tx.unbounded_send(cancel) + .expect("cancellation has a separate control lane"); + assert!(rx.next().now_or_never().unwrap().unwrap().is_urgent()); + } - let ((), (), ()) = futures::try_join!( - Channel { - rx: client_channel.rx, - tx: self.tx, - } - .copy(), - Channel { - rx: self.rx, - tx: client_channel.tx, + #[test] + fn cancellation_admission_is_urgent_at_every_queue_occupancy() { + use super::admission::ReceiverClose as _; + + for queued in 0..=3 { + for asynchronous in [false, true] { + for wrapped in [false, true] { + let (channel, _peer) = Channel::duplex_with_limits(ConnectionLimits { + max_queued_frames: 3, + ..ConnectionLimits::default() + }); + let (tx, mut rx) = application_channel(channel.tx.admission()); + for _ in 0..queued { + tx.unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("data", serde_json::json!({})).unwrap(), + }) + .unwrap(); + } + let params = serde_json::json!({"requestId":"waiting"}); + let cancel = if wrapped { + UntypedMessage::new( + "_proxy/successor", + serde_json::json!({"method":"$/cancel_request","params":params}), + ) + } else { + UntypedMessage::new("$/cancel_request", params) + } + .unwrap(); + let message = OutgoingMessage::Notification { untyped: cancel }; + if asynchronous { + tx.send(message) + .now_or_never() + .expect("urgent admission does not wait for ordinary queue space") + .unwrap(); + } else { + tx.unbounded_send(message).unwrap(); + } + let urgent = future::poll_fn(|cx| rx.poll_urgent(cx)) + .now_or_never() + .expect("readiness gate must observe cancellation") + .unwrap(); + assert!(urgent.is_urgent()); + for _ in 0..queued { + assert!(!rx.next().now_or_never().unwrap().unwrap().is_urgent()); + } + assert!(rx.next().now_or_never().is_none()); + } } - .copy(), - client_future, - )?; - Ok(()) + } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { - (self, Box::pin(future::ready(Ok(())))) + #[test] + fn cancellation_passes_waiting_request_and_saturated_data_lane() { + for capacity in [1, 3, ConnectionLimits::default().max_queued_frames] { + check_cancellation_passes_waiting_request(capacity); + } } -} -#[cfg(test)] -mod tests { - use super::*; + fn check_cancellation_passes_waiting_request(capacity: usize) { + let (transport, mut peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 1024, + max_queued_bytes: 4096, + max_queued_frames: capacity, + }); + let (message_tx, message_rx) = application_channel(transport.tx.admission()); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_tx, _dynamic_rx) = admission::channel(); + let pending_replies = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx.clone(), + task_tx, + dynamic_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending_replies.registrar(), + ProtocolMode::disabled(), + ); + let (ready_tx, ready_rx) = oneshot::channel(); + let sent = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("not-ready", serde_json::json!({})).unwrap(), + async move { ready_rx.await.map_err(crate::Error::into_internal_error) }, + ); + let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( + message_rx, + pending_replies, + transport.tx, + ProtocolCompat::new(ProtocolMode::disabled()), + IncomingClosed::new(), + )); + assert!(actor.as_mut().now_or_never().is_none()); + message_tx + .unbounded_send(OutgoingMessage::Notification { + untyped: UntypedMessage::new("data", serde_json::json!({})).unwrap(), + }) + .unwrap(); + connection + .send_cancel_request(sent.id().clone()) + .expect("urgent lane should remain available"); + assert!(actor.as_mut().now_or_never().is_none()); + let error = sent + .block_task() + .now_or_never() + .expect("cancel settles without readiness") + .expect_err("unpublished request must be cancelled locally"); + assert_eq!(error.code, crate::ErrorCode::RequestCancelled); + assert!( + ready_tx.send(()).is_err(), + "cancelled readiness future must be dropped" + ); + let data = peer.rx.next().now_or_never().unwrap().unwrap(); + assert!( + matches!(data.frame(), TransportFrame::Single(RawJsonRpcMessage::Notification(n)) + if n.method.as_ref() == "data") + ); + assert!( + peer.rx.next().now_or_never().is_none(), + "never publish a request after its cancellation" + ); + } + + #[test] + fn cancellation_of_queued_request_does_not_wait_for_unrelated_readiness() { + for capacity in [1, 3, ConnectionLimits::default().max_queued_frames] { + check_cancellation_of_queued_request(capacity); + } + } + + fn check_cancellation_of_queued_request(capacity: usize) { + let (transport, mut peer) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 1024, + max_queued_bytes: 8192, + max_queued_frames: capacity, + }); + let (message_tx, message_rx) = application_channel(transport.tx.admission()); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_tx, _dynamic_rx) = admission::channel(); + let pending_replies = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx, + task_tx, + dynamic_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending_replies.registrar(), + ProtocolMode::disabled(), + ); + let (ready_tx, ready_rx) = oneshot::channel(); + let first = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("first", serde_json::json!({})).unwrap(), + async move { ready_rx.await.map_err(crate::Error::into_internal_error) }, + ); + let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( + message_rx, + pending_replies, + transport.tx, + ProtocolCompat::new(ProtocolMode::disabled()), + IncomingClosed::new(), + )); + assert!(actor.as_mut().now_or_never().is_none()); + let second = connection.send_request_to( + crate::role::UntypedRole, + UntypedMessage::new("second", serde_json::json!({})).unwrap(), + ); + second.cancel().unwrap(); + assert!(actor.as_mut().now_or_never().is_none()); + let error = second + .block_task() + .now_or_never() + .expect("queued cancellation cannot wait for first") + .expect_err("second request was never published"); + assert_eq!(error.code, crate::ErrorCode::RequestCancelled); + assert!(peer.rx.next().now_or_never().is_none()); + + ready_tx.send(()).unwrap(); + assert!(actor.as_mut().now_or_never().is_none()); + let frame = peer.rx.next().now_or_never().unwrap().unwrap(); + assert!( + matches!(frame.frame(), TransportFrame::Single(RawJsonRpcMessage::Request(r)) + if r.method.as_ref() == "first") + ); + drop(frame); + assert!( + peer.rx.next().now_or_never().is_none(), + "second must not run later" + ); + first.detach(); + } + + #[cfg(feature = "unstable_mcp_over_acp")] + #[test] + fn protected_operation_tracking_is_bounded_across_sequential_completions() { + let (message_tx, _message_rx) = admission::channel(); + let (task_tx, task_rx) = task_actor::task_channel(2); + let (dynamic_tx, _dynamic_rx) = admission::channel(); + let pending_replies = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx, + task_tx, + dynamic_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending_replies.registrar(), + ProtocolMode::disabled(), + ); + let mut driver = Box::pin(task_actor::task_actor(task_rx, &connection, 2)); + for _ in 0..1000 { + connection.spawn_protected(async { Ok(()) }).unwrap(); + assert_eq!( + connection + .protected_operations + .lock() + .unwrap() + .pending + .len(), + 1, + "previous completions must be reaped before another admission" + ); + assert!(driver.as_mut().now_or_never().is_none()); + } + assert!( + connection + .wait_protected_operations() + .now_or_never() + .is_some() + ); + assert!( + connection + .protected_operations + .lock() + .unwrap() + .pending + .is_empty() + ); + assert!(connection.spawn_protected(async { Ok(()) }).is_err()); + } #[cfg(feature = "unstable_protocol_v2")] fn connection_with_task_receiver() -> ( ConnectionTo, - mpsc::UnboundedReceiver, + admission::SimpleReceiver, ) { - let (message_tx, _message_rx) = mpsc::unbounded(); - let (task_tx, task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, _message_rx) = admission::channel(); + let (task_tx, task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -6708,9 +9065,9 @@ mod tests { #[cfg(feature = "unstable_protocol_v2")] #[test] fn v2_proxy_rejects_explicitly_prewrapped_initialize_request() { - let (message_tx, message_rx) = mpsc::unbounded(); - let (task_tx, _task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, message_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -6734,12 +9091,21 @@ mod tests { }; let sent = connection.send_request_to(Agent, request); - let (transport_tx, mut transport_rx) = mpsc::unbounded(); + let ( + Channel { + tx: transport_tx, .. + }, + Channel { + rx: mut transport_rx, + .. + }, + ) = Channel::duplex(); let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( message_rx, pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::v2_proxy()), + IncomingClosed::new(), )); assert!( actor.as_mut().now_or_never().is_none(), @@ -6859,11 +9225,11 @@ mod tests { fn connection_with_dynamic_handler_receiver() -> ( ConnectionTo, - mpsc::UnboundedReceiver>, + admission::SimpleReceiver>, ) { - let (message_tx, _message_rx) = mpsc::unbounded(); - let (task_tx, _task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, _message_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_handler_tx, dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -6900,12 +9266,12 @@ mod tests { fn connection_for_response_hook_tests() -> ( ConnectionTo, - mpsc::UnboundedReceiver, + admission::SimpleReceiver, PendingReplies, ) { - let (message_tx, message_rx) = mpsc::unbounded(); - let (task_tx, _task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, message_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -6925,6 +9291,138 @@ mod tests { ) } + fn budgeted_request_connection( + limits: ConnectionLimits, + ) -> ( + ConnectionTo, + admission::Receiver, + PendingReplies, + FrameAdmission, + ) { + let (channel, _) = Channel::duplex_with_limits(limits); + let admission = channel.tx.admission(); + let (message_tx, message_rx) = application_channel(admission.clone()); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); + let pending = PendingReplies::default(); + let connection = ConnectionTo::new( + crate::role::UntypedRole, + message_tx, + task_tx, + dynamic_handler_tx, + future::ready(Ok::<(), crate::Error>(())).boxed().shared(), + pending.registrar(), + ProtocolMode::disabled(), + ); + (connection, message_rx, pending, admission) + } + + #[test] + fn pending_request_metadata_remains_charged_after_queue_consumption() { + let (connection, mut rx, pending, admission) = + budgeted_request_connection(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 1312, + max_queued_frames: 32, + }); + let method = "m".repeat(140); + let request = || UntypedMessage::new(&method, serde_json::json!({})).unwrap(); + let first = connection.send_request_to(crate::role::UntypedRole, request()); + let first_id = first.id().clone(); + let queued = futures::FutureExt::now_or_never(futures::StreamExt::next(&mut rx)) + .unwrap() + .unwrap(); + drop(queued); + assert!(admission.0.state.lock().unwrap().used > 0); + let second = connection.send_request_to(crate::role::UntypedRole, request()); + assert!(futures::executor::block_on(second.block_task()).is_err()); + assert!(pending.remove(&first_id).is_some()); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + drop(first); + } + + #[test] + fn rejected_admitted_request_fails_without_leaking_pending_reply() { + let (connection, mut rx, pending, admission) = + budgeted_request_connection(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 2048, + max_queued_frames: 1, + }); + admission::ReceiverClose::close(&mut rx); + let sent = connection.send_request_to( + crate::role::UntypedRole, + UntypedMessage::new("rejected", serde_json::json!({})).unwrap(), + ); + assert!(!pending.contains(sent.id())); + assert!(futures::executor::block_on(sent.block_task()).is_err()); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + + #[test] + fn cancelling_request_releases_pending_metadata_after_queue_consumption() { + let (connection, mut rx, pending, admission) = + budgeted_request_connection(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 1312, + max_queued_frames: 8, + }); + let sent = connection.send_request_to( + crate::role::UntypedRole, + UntypedMessage::new(&"m".repeat(140), serde_json::json!({})).unwrap(), + ); + let id = sent.id().clone(); + drop(futures::FutureExt::now_or_never(futures::StreamExt::next(&mut rx)).unwrap()); + assert!(pending.contains(&id)); + drop(sent); + assert!(!pending.contains(&id)); + // Drop the cancellation notification too; no payload remains admitted. + drop(futures::FutureExt::now_or_never(futures::StreamExt::next(&mut rx)).unwrap()); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + + #[test] + fn incoming_eof_releases_many_pending_method_charges_after_error_consumption() { + let (channel, _) = Channel::duplex_with_limits(ConnectionLimits { + max_frame_bytes: 512, + max_queued_bytes: 2100, + max_queued_frames: 32, + }); + let admission = channel.tx.admission(); + let pending = PendingReplies::with_capacity(32); + let mut receivers = Vec::new(); + for i in 0..7 { + let method = "m".repeat(120); + let id = RequestId::Str(format!("{i:036}")); + let charge = admission + .try_reserve_bytes(method.len() + 36 + 64, true) + .unwrap(); + let (sender, receiver) = oneshot::channel(); + assert!(pending.registrar().subscribe( + id, + PendingReply { + method, + metadata_bytes: Some(charge), + role_id: crate::role::UntypedRole.role_id(), + sender, + cancellation_disarm: SentRequestCancellationDisarm::new(), + ordering: ResponseOrdering::default(), + response_route_hook: None, + }, + &IncomingClosed::new(), + )); + receivers.push(receiver); + } + assert!(admission.try_reserve_bytes(220, true).is_none()); + assert_eq!(pending.close_incoming(), 7); + assert!( + admission.0.state.lock().unwrap().used > 0, + "failed results still own their method text" + ); + drop(receivers); + assert_eq!(admission.0.state.lock().unwrap().used, 0); + } + #[cfg(feature = "unstable_protocol_v2")] fn route_test_response( request_id: RequestId, @@ -6935,7 +9433,7 @@ mod tests { .remove(&request_id) .expect("the request should have a pending reply"); let (dispatch, _) = - incoming_actor::dispatch_from_response(request_id, pending_reply, result); + incoming_actor::dispatch_from_response(request_id, pending_reply, result, None); let Dispatch::Response(result, router) = dispatch else { panic!("expected a response dispatch"); }; @@ -7068,12 +9566,21 @@ mod tests { async move { ready_rx.await.map_err(crate::Error::into_internal_error) }, ); - let (transport_tx, mut transport_rx) = mpsc::unbounded(); + let ( + Channel { + tx: transport_tx, .. + }, + Channel { + rx: mut transport_rx, + .. + }, + ) = Channel::duplex(); let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( message_rx, pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::disabled()), + IncomingClosed::new(), )); assert!( @@ -7098,13 +9605,43 @@ mod tests { .expect("the ready request should be published") .expect("the transport queue should remain open"); assert!(matches!( - frame, + frame.frame(), TransportFrame::Single(RawJsonRpcMessage::Request(_)) )); drop(sent); } + #[test] + fn pending_outgoing_readiness_is_cancelled_on_shutdown() { + let (connection, message_rx, pending_replies) = connection_for_response_hook_tests(); + let sent = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("waiting", serde_json::json!({})).unwrap(), + future::pending::>(), + ); + let ( + Channel { + tx, + rx: mut transport_rx, + }, + _peer, + ) = Channel::duplex(); + let shutdown = IncomingClosed::new(); + let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( + message_rx, + pending_replies, + tx, + ProtocolCompat::new(ProtocolMode::disabled()), + shutdown.clone(), + )); + assert!(actor.as_mut().now_or_never().is_none()); + shutdown.begin_close(); + assert!(actor.as_mut().now_or_never().is_none()); + assert!(transport_rx.next().now_or_never().is_none()); + assert!(futures::executor::block_on(sent.block_task()).is_err()); + } + #[test] fn ordered_blocking_transform_precedes_response_acknowledgment() { let (connection, _message_rx, pending_replies) = connection_for_response_hook_tests(); @@ -7121,6 +9658,7 @@ mod tests { request_id, pending_reply, Err(crate::Error::invalid_params()), + None, ); let Dispatch::Response(result, router) = dispatch else { panic!("expected a response dispatch"); @@ -7180,12 +9718,21 @@ mod tests { }, ); - let (transport_tx, mut transport_rx) = mpsc::unbounded(); + let ( + Channel { + tx: transport_tx, .. + }, + Channel { + rx: mut transport_rx, + .. + }, + ) = Channel::duplex(); let mut actor = Box::pin(outgoing_actor::outgoing_protocol_actor( message_rx, pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::disabled()), + IncomingClosed::new(), )); assert!( @@ -7205,9 +9752,9 @@ mod tests { #[test] fn ordered_request_is_marked_before_entering_outgoing_queue() { - let (message_tx, mut message_rx) = mpsc::unbounded(); - let (task_tx, mut task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, _dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, mut message_rx) = admission::channel(); + let (task_tx, mut task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_handler_tx, _dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); @@ -7250,6 +9797,7 @@ mod tests { request_id, pending_reply, Ok(serde_json::json!({"ok": true})), + None, ); let Dispatch::Response(result, router) = dispatch else { panic!("expected a response dispatch"); @@ -7282,7 +9830,7 @@ mod tests { } fn next_dynamic_handler_message( - receiver: &mut mpsc::UnboundedReceiver>, + receiver: &mut (impl futures::Stream> + Unpin), ) -> Option> { futures::FutureExt::now_or_never(futures::StreamExt::next(receiver)) .expect("dynamic-handler receiver should be ready") @@ -7291,9 +9839,9 @@ mod tests { #[cfg(feature = "unstable_protocol_v2")] #[test] fn v2_dynamic_handler_guard_registers_and_removes_handler() { - let (message_tx, _message_rx) = mpsc::unbounded(); - let (task_tx, _task_rx) = mpsc::unbounded(); - let (dynamic_handler_tx, mut dynamic_handler_rx) = mpsc::unbounded(); + let (message_tx, _message_rx) = admission::channel(); + let (task_tx, _task_rx) = task_actor::task_channel(admission::QUEUE_CAPACITY); + let (dynamic_handler_tx, mut dynamic_handler_rx) = admission::channel(); let transport_completion: SharedTransportCompletion = future::ready(Ok::<(), crate::Error>(())).boxed().shared(); let pending_replies = PendingReplies::default(); diff --git a/src/agent-client-protocol/src/jsonrpc/admission.rs b/src/agent-client-protocol/src/jsonrpc/admission.rs new file mode 100644 index 00000000..4c5d8a1d --- /dev/null +++ b/src/agent-client-protocol/src/jsonrpc/admission.rs @@ -0,0 +1,278 @@ +//! Finite queues for synchronously invoked dispatcher APIs. +//! +//! Dispatch callbacks cannot await capacity: the receiver may depend on that +//! callback returning. External producers can await `send` instead. +use futures::Stream; +use std::pin::Pin; +use std::sync::Arc; +use std::task::{Context, Poll}; + +use super::{FrameAdmission, FramePermit}; + +pub const QUEUE_CAPACITY: usize = 32; + +pub struct Sender { + inner: Arc>, +} + +struct SenderInner { + tx: async_channel::Sender, + urgent_tx: Option>, + admission: Option>, + capacity: usize, +} + +struct Admission { + budget: FrameAdmission, + measure: fn(&T) -> Result, + attach: fn(T, FramePermit) -> T, + control: fn(&T) -> bool, + urgent: fn(&T) -> bool, +} + +impl Clone for Sender { + fn clone(&self) -> Self { + Self { + inner: self.inner.clone(), + } + } +} + +impl std::fmt::Debug for Sender { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("AdmissionSender").finish_non_exhaustive() + } +} + +impl Sender { + pub fn byte_admission(&self) -> Option { + self.inner + .admission + .as_ref() + .map(|admission| admission.budget.clone()) + } + + pub fn queue_capacity(&self) -> usize { + self.inner.capacity + } + + pub fn unbounded_send(&self, item: T) -> Result<(), SendError> { + // A readiness-blocked consumer polls only the urgent lane, regardless + // of ordinary queue occupancy. The outgoing actor settles cancellation + // locally if its request has not yet been published. + let urgent = self + .inner + .admission + .as_ref() + .is_some_and(|admission| (admission.urgent)(&item)); + let item = if let Some(admission) = &self.inner.admission { + let bytes = (admission.measure)(&item).map_err(|error| SendError { + item: None, + reason: error.to_string(), + }); + // Preserve ownership of the rejected message even when sizing fails. + let bytes = match bytes { + Ok(bytes) => bytes, + Err(error) => { + return Err(SendError { + item: Some(item), + ..error + }); + } + }; + let permit = admission + .budget + .try_reserve_bytes(bytes, !(admission.control)(&item)) + .ok_or_else(|| SendError { + item: None, + reason: "outgoing application byte capacity exceeded".into(), + }); + let permit = match permit { + Ok(permit) => permit, + Err(error) => { + return Err(SendError { + item: Some(item), + ..error + }); + } + }; + (admission.attach)(item, permit) + } else { + item + }; + let tx = if urgent { + self.inner.urgent_tx.as_ref().expect("urgent lane exists") + } else { + &self.inner.tx + }; + tx.try_send(item).map_err(|error| SendError { + item: Some(error.into_inner()), + reason: "outgoing application queue full or closed".into(), + }) + } + + pub async fn send(&self, item: T) -> Result<(), crate::Error> { + let urgent = self + .inner + .admission + .as_ref() + .is_some_and(|admission| (admission.urgent)(&item)); + let tx = if urgent { + self.inner.urgent_tx.as_ref().expect("urgent lane exists") + } else { + &self.inner.tx + }; + let item = if let Some(admission) = &self.inner.admission { + let bytes = (admission.measure)(&item)?; + let reserve = admission + .budget + .reserve_bytes(bytes, !(admission.control)(&item)); + let permit = + match futures::future::select(Box::pin(reserve), Box::pin(tx.closed())).await { + futures::future::Either::Left((permit, _)) => permit?, + futures::future::Either::Right(_) => { + return Err(crate::util::internal_error( + "outgoing application queue closed", + )); + } + }; + (admission.attach)(item, permit) + } else { + item + }; + tx.send(item).await.map_err(crate::util::internal_error) + } +} + +#[derive(Debug)] +pub struct SendError { + item: Option, + reason: String, +} + +impl SendError { + pub fn into_inner(self) -> T { + self.item.expect("send errors retain their rejected item") + } +} + +impl std::fmt::Display for SendError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(&self.reason) + } +} + +impl std::error::Error for SendError {} + +#[cfg(test)] +pub fn channel() -> (Sender, SimpleReceiver) { + channel_with_capacity(QUEUE_CAPACITY) +} + +pub fn channel_with_capacity(capacity: usize) -> (Sender, SimpleReceiver) { + let capacity = capacity.max(1); + let (tx, rx) = async_channel::bounded(capacity); + ( + Sender { + inner: Arc::new(SenderInner { + tx, + urgent_tx: None, + admission: None, + capacity, + }), + }, + SimpleReceiver(Box::pin(rx)), + ) +} + +pub struct SimpleReceiver(Pin>>); + +impl Stream for SimpleReceiver { + type Item = T; + + fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + self.0.as_mut().poll_next(cx) + } +} + +pub(super) trait ReceiverClose: Stream { + fn close(&mut self); + fn poll_urgent(&mut self, _cx: &mut Context<'_>) -> Poll> { + Poll::Pending + } +} + +impl ReceiverClose for SimpleReceiver { + fn close(&mut self) { + self.0.close(); + } +} + +pub struct Receiver { + normal: SimpleReceiver, + urgent: SimpleReceiver, +} + +impl Stream for Receiver { + type Item = T; + + fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let this = self.get_mut(); + let urgent_closed = match Pin::new(&mut this.urgent).poll_next(cx) { + Poll::Ready(Some(item)) => return Poll::Ready(Some(item)), + Poll::Ready(None) => true, + Poll::Pending => false, + }; + match Pin::new(&mut this.normal).poll_next(cx) { + Poll::Ready(Some(item)) => Poll::Ready(Some(item)), + Poll::Ready(None) if urgent_closed => Poll::Ready(None), + _ => Poll::Pending, + } + } +} + +impl ReceiverClose for Receiver { + fn close(&mut self) { + self.normal.close(); + self.urgent.close(); + } + + fn poll_urgent(&mut self, cx: &mut Context<'_>) -> Poll> { + match Pin::new(&mut self.urgent).poll_next(cx) { + Poll::Ready(None) => Poll::Pending, + result => result, + } + } +} + +pub fn budgeted_channel( + admission: FrameAdmission, + measure: fn(&T) -> Result, + attach: fn(T, FramePermit) -> T, + control: fn(&T) -> bool, + urgent: fn(&T) -> bool, +) -> (Sender, Receiver) { + let capacity = admission.limits().max_queued_frames.max(1); + let (tx, rx) = async_channel::bounded(capacity); + let (urgent_tx, urgent_rx) = async_channel::bounded(capacity); + ( + Sender { + inner: Arc::new(SenderInner { + tx, + urgent_tx: Some(urgent_tx), + admission: Some(Admission { + budget: admission, + measure, + attach, + control, + urgent, + }), + capacity, + }), + }, + Receiver { + normal: SimpleReceiver(Box::pin(rx)), + urgent: SimpleReceiver(Box::pin(urgent_rx)), + }, + ) +} diff --git a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs index 35383168..0fe5cdb6 100644 --- a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs @@ -1,6 +1,5 @@ // Types re-exported from crate root use futures::StreamExt as _; -use futures::channel::mpsc; use futures::stream; use futures_concurrency::stream::StreamExt as _; use rustc_hash::FxHashMap; @@ -18,21 +17,22 @@ use crate::jsonrpc::PendingReplies; use crate::jsonrpc::PendingReply; use crate::jsonrpc::RawJsonRpcMessage; use crate::jsonrpc::RawJsonRpcParams; +use crate::jsonrpc::RawJsonRpcResponse as Response; use crate::jsonrpc::RequestReplyTarget; use crate::jsonrpc::Responder; use crate::jsonrpc::ResponseDestination; use crate::jsonrpc::ResponseDispatch; use crate::jsonrpc::ResponseRouter; use crate::jsonrpc::TransportBatchEntry; -use crate::jsonrpc::TransportFrame; use crate::jsonrpc::dynamic_handler::DynHandleDispatchFrom; use crate::jsonrpc::dynamic_handler::DynamicHandlerMessage; use crate::jsonrpc::outgoing_actor::send_raw_message; use crate::jsonrpc::protocol_compat::ProtocolCompat; +use crate::jsonrpc::{BudgetedFrame, FramePermit, TransportFrame}; use crate::jsonrpc::{is_response_only_shape, raw_is_response_only_shape}; use crate::role::Role; -use crate::schema::v1::{RequestId, Response}; +use crate::schema::v1::RequestId; use super::Handled; @@ -59,8 +59,8 @@ impl IncomingHandlers { pub(super) async fn incoming_protocol_actor( counterpart: Counterpart, connection: &ConnectionTo, - transport_rx: mpsc::UnboundedReceiver, - dynamic_handler_rx: mpsc::UnboundedReceiver>, + transport_rx: super::FrameReceiver, + dynamic_handler_rx: super::admission::SimpleReceiver>, pending_replies: PendingReplies, handlers: IncomingHandlers< impl HandleDispatchFrom, @@ -85,7 +85,7 @@ pub(super) async fn incoming_protocol_actor( let mut dynamic_handlers: FxHashMap>> = FxHashMap::default(); - let mut pending_messages: Vec = vec![]; + let mut pending_messages: Vec = vec![]; let request_cancellations = super::RequestCancellationRegistry::new(); let mut on_close = Some(on_close); @@ -128,6 +128,7 @@ pub(super) async fn incoming_protocol_actor( } IncomingProtocolMsg::Transport(frame) => { + let (frame, permit) = frame.into_parts(); let (entries, batch_completion) = frame_entries(frame); for (message, destination) in entries { match message { @@ -156,6 +157,7 @@ pub(super) async fn incoming_protocol_actor( &mut handler, &mut pending_messages, &request_cancellations, + permit.clone(), ) .await?; } @@ -196,6 +198,7 @@ pub(super) async fn incoming_protocol_actor( &mut handler, &mut pending_messages, &request_cancellations, + permit.clone(), ) .await?; } @@ -208,15 +211,19 @@ pub(super) async fn incoming_protocol_actor( Ok(RawJsonRpcMessage::Response(response)) => { let (id, result) = match response { Response::Result { id, result } => (id, Ok(result)), - Response::Error { id, error } => (id, Err(error)), + Response::Error { id, error } => (id, Err(error.into_acp_error())), }; tracing::trace!(?id, "Handling response"); if let Some(pending_reply) = pending_replies.remove(&id) { let result = protocol_compat .incoming_response(&pending_reply.method, result); - let (dispatch, response_dispatch) = - dispatch_from_response(id, pending_reply, result); + let (dispatch, response_dispatch) = dispatch_from_response( + id, + pending_reply, + result, + Some(permit.clone()), + ); dispatch_dispatch( counterpart.clone(), connection, @@ -225,6 +232,7 @@ pub(super) async fn incoming_protocol_actor( &mut handler, &mut pending_messages, &request_cancellations, + permit.clone(), ) .await?; if let Some(ack_rx) = response_dispatch.complete() { @@ -260,6 +268,13 @@ pub(super) async fn incoming_protocol_actor( } message @ (IncomingProtocolMsg::Transport(_) | IncomingProtocolMsg::TransportClosed) => { + if queued_transport_messages.len() + >= connection.message_tx.queue_capacity() + { + return Err(crate::util::internal_error( + "transport frames exceed barrier queue capacity", + )); + } queued_transport_messages.push_back(message); } } @@ -303,14 +318,18 @@ async fn handle_dynamic_handler_message( message: DynamicHandlerMessage, connection: &ConnectionTo, dynamic_handlers: &mut FxHashMap>>, - pending_messages: &mut Vec, + pending_messages: &mut Vec, ) -> Result<(), crate::Error> { match message { DynamicHandlerMessage::AddDynamicHandler(uuid, mut handler) => { // Before adding the new handler, give it a chance to process // any pending messages. let mut new_pending_messages = vec![]; - for pending_message in std::mem::take(pending_messages) { + for DeferredDispatch { + dispatch: pending_message, + permit, + } in std::mem::take(pending_messages) + { tracing::trace!(method = pending_message.method(), handler = ?handler.dyn_describe_chain(), "Retrying message"); let reply_target = pending_message.handler_error_target(); let handler_attempt = reply_target @@ -329,7 +348,10 @@ async fn handle_dynamic_handler_message( retry: _, }) => { tracing::trace!(method = m.method(), handler = ?handler.dyn_describe_chain(), "Message not handled"); - new_pending_messages.push(m); + new_pending_messages.push(DeferredDispatch { + dispatch: m, + permit: permit.clone(), + }); } Err(err) => { tracing::warn!(?err, handler = ?handler.dyn_describe_chain(), "Dynamic handler errored on pending message"); @@ -341,6 +363,13 @@ async fn handle_dynamic_handler_message( *pending_messages = new_pending_messages; // Add handler so it will be used for future incoming messages. + if !dynamic_handlers.contains_key(&uuid) + && dynamic_handlers.len() >= connection.message_tx.queue_capacity() + { + return Err(crate::util::internal_error( + "dynamic handler capacity exceeded", + )); + } dynamic_handlers.insert(uuid, handler); } DynamicHandlerMessage::RemoveDynamicHandler(uuid) => { @@ -356,11 +385,16 @@ async fn handle_dynamic_handler_message( #[derive(Debug)] enum IncomingProtocolMsg { - Transport(TransportFrame), + Transport(BudgetedFrame), TransportClosed, DynamicHandler(DynamicHandlerMessage), } +struct DeferredDispatch { + dispatch: Dispatch, + permit: FramePermit, +} + fn frame_entries( frame: TransportFrame, ) -> ( @@ -472,11 +506,17 @@ pub(super) fn dispatch_from_response( id: RequestId, pending_reply: PendingReply, result: Result, + frame_bytes: Option, ) -> (Dispatch, ResponseDispatch) { let response_dispatch = ResponseDispatch::default(); // Create a Dispatch::Response with a ResponseRouter that routes to the oneshot - let router = ResponseRouter::new(id.clone(), pending_reply, response_dispatch.clone()); + let router = ResponseRouter::new( + id.clone(), + pending_reply, + response_dispatch.clone(), + frame_bytes, + ); (Dispatch::Response(result, router), response_dispatch) } @@ -485,14 +525,19 @@ pub(super) fn dispatch_from_response( fields(method = dispatch.method()), level = "trace", )] +#[expect( + clippy::too_many_arguments, + reason = "one dispatch carries its retained frame admission" +)] async fn dispatch_dispatch( counterpart: Counterpart, connection: &ConnectionTo, mut dispatch: Dispatch, dynamic_handlers: &mut FxHashMap>>, handler: &mut impl HandleDispatchFrom, - pending_messages: &mut Vec, + pending_messages: &mut Vec, request_cancellations: &super::RequestCancellationRegistry, + permit: FramePermit, ) -> Result<(), crate::Error> { tracing::trace!(?dispatch, "dispatch_dispatch"); @@ -607,7 +652,15 @@ async fn dispatch_dispatch( ?method, "Retrying message as new dynamic handlers are added" ); - pending_messages.push(dispatch); + if pending_messages.len() >= connection.message_tx.queue_capacity() { + return handle_handler_error( + connection, + error_target, + method, + crate::util::internal_error("pending dispatch capacity exceeded"), + ); + } + pending_messages.push(DeferredDispatch { dispatch, permit }); Ok(()) } else { match dispatch { @@ -617,6 +670,13 @@ async fn dispatch_dispatch( } Dispatch::Request(_, responder) => { tracing::info!(?method, "Rejecting request with error, no handler"); + #[cfg(feature = "unstable_mcp_over_acp")] + if method == "mcp/message" { + return responder.respond_with_error(crate::Error::new( + crate::mcp_server::MCP_SERVER_UNAVAILABLE, + "MCP server unavailable", + )); + } responder.respond_with_error(crate::Error::method_not_found().data(method)) } Dispatch::Response(result, router) => { diff --git a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs index effe164c..f184b847 100644 --- a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs @@ -1,12 +1,15 @@ // Types re-exported from crate root use futures::StreamExt as _; -use futures::channel::mpsc; +use futures::future; +use std::task::Poll; use crate::jsonrpc::protocol_compat::ProtocolCompat; -use crate::jsonrpc::{OutgoingMessage, PendingReplies, RawJsonRpcMessage, TransportFrame}; +use crate::jsonrpc::{ + FramePermit, OutgoingMessage, PendingReplies, RawJsonRpcMessage, TransportFrame, UntypedMessage, +}; use crate::schema::v1::RequestId; -pub type OutgoingMessageTx = mpsc::UnboundedSender; +pub type OutgoingMessageTx = super::admission::Sender; pub(crate) fn send_raw_message( tx: &OutgoingMessageTx, @@ -17,6 +20,45 @@ pub(crate) fn send_raw_message( .map_err(crate::util::internal_error) } +async fn publish( + tx: &super::FrameSender, + frame: TransportFrame, + permit: Option, +) -> Result<(), crate::Error> { + match permit { + Some(permit) => tx.send_admitted(frame, permit).await, + None => tx.send_frame(frame).await, + } + .map_err(crate::Error::into_internal_error) +} + +async fn publish_notification( + tx: &super::FrameSender, + protocol_compat: &ProtocolCompat, + pending_replies: &PendingReplies, + untyped: UntypedMessage, + permit: Option, +) -> Result<(), crate::Error> { + if let Some(id) = super::outgoing_cancellation_id(&untyped) + && pending_replies.cancel_unpublished(&id) + { + return Ok(()); + } + let messages = protocol_compat.outgoing_notification(untyped)?; + // ProtocolCompat currently emits exactly one notification. A future + // expansion needs separately admitted charges for each additional output. + if messages.len() > 1 { + return Err(crate::util::internal_error( + "notification expansion exceeds application admission", + )); + } + if let Some(untyped) = messages.into_iter().next() { + let message = untyped.into_raw_jsonrpc_message(None)?; + publish(tx, TransportFrame::Single(message), permit).await?; + } + Ok(()) +} + /// Outgoing protocol actor: Converts application-level OutgoingMessage to protocol-level RawJsonRpcMessage. /// /// This actor handles JSON-RPC protocol semantics: @@ -25,15 +67,20 @@ pub(crate) fn send_raw_message( /// /// This is the protocol layer - it has no knowledge of how messages are transported. pub(super) async fn outgoing_protocol_actor( - mut outgoing_rx: mpsc::UnboundedReceiver, + mut outgoing_rx: impl Unpin + super::admission::ReceiverClose, pending_replies: PendingReplies, - transport_tx: mpsc::UnboundedSender, + transport_tx: super::FrameSender, protocol_compat: ProtocolCompat, + shutdown: super::IncomingClosed, ) -> Result<(), crate::Error> { let mut drain_waiters = Vec::new(); while let Some(message) = outgoing_rx.next().await { tracing::debug!(?message, "outgoing_protocol_actor"); + let (message, permit) = match message { + OutgoingMessage::Admitted { message, permit } => (*message, Some(permit)), + message => (message, None), + }; // Create the message to be sent over the transport let (json_rpc_message, destination) = match message { @@ -45,18 +92,14 @@ pub(super) async fn outgoing_protocol_actor( continue; } OutgoingMessage::BatchDispatchComplete { completion } => { - if let Some(frame) = completion.complete() { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some((frame, permit)) = completion.complete_admitted(permit) { + publish(&transport_tx, frame, permit).await?; } continue; } OutgoingMessage::BatchHandlerAttemptComplete { destination } => { - if let Some(frame) = destination.finish_handler_attempt() { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some((frame, permit)) = destination.finish_handler_attempt_admitted(permit) { + publish(&transport_tx, frame, permit).await?; } continue; } @@ -78,10 +121,8 @@ pub(super) async fn outgoing_protocol_actor( ))), ); let fallback = RawJsonRpcMessage::response(id, fallback); - if let Some(frame) = destination.abandon(fallback) { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some((frame, permit)) = destination.abandon_admitted(fallback, permit) { + publish(&transport_tx, frame, permit).await?; } continue; } @@ -99,19 +140,70 @@ pub(super) async fn outgoing_protocol_actor( continue; } - if let Some(readiness) = readiness - && let Err(error) = readiness.await - { - tracing::warn!( - ?id, - %method, - ?error, - "Outgoing request readiness failed" - ); - if let Some(pending_reply) = pending_replies.remove(&id) { - pending_reply.fail(error); + if let Some(readiness) = readiness { + enum Gate { + Ready(Result<(), crate::Error>), + Shutdown, + Urgent(OutgoingMessage), + } + let mut readiness = Box::pin(readiness); + let mut closing = Box::pin(shutdown.shutdown_requested()); + let mut skip_request = false; + loop { + let gate = future::poll_fn(|cx| { + if let Poll::Ready(result) = readiness.as_mut().poll(cx) { + return Poll::Ready(Gate::Ready(result)); + } + if closing.as_mut().poll(cx).is_ready() { + return Poll::Ready(Gate::Shutdown); + } + match outgoing_rx.poll_urgent(cx) { + Poll::Ready(Some(message)) => Poll::Ready(Gate::Urgent(message)), + _ => Poll::Pending, + } + }) + .await; + match gate { + Gate::Ready(Ok(())) => break, + Gate::Ready(Err(error)) => { + tracing::warn!(?id, %method, ?error, "Outgoing request readiness failed"); + if let Some(pending_reply) = pending_replies.remove(&id) { + pending_reply.fail(error); + } + skip_request = true; + break; + } + Gate::Shutdown => { + if let Some(pending_reply) = pending_replies.remove(&id) { + pending_reply.fail(crate::util::internal_error("connection shut down while waiting for outgoing request readiness")); + } + skip_request = true; + break; + } + Gate::Urgent(OutgoingMessage::Admitted { message, permit }) => { + if let OutgoingMessage::Notification { untyped } = *message { + publish_notification( + &transport_tx, + &protocol_compat, + &pending_replies, + untyped, + Some(permit), + ) + .await?; + } + if !pending_replies.contains(&id) { + skip_request = true; + break; + } + } + Gate::Urgent(_) => unreachable!( + "urgent admission only accepts cancellation notifications" + ), + } + } + if skip_request { + continue; } - continue; } if !pending_replies.contains(&id) { @@ -133,11 +225,13 @@ pub(super) async fn outgoing_protocol_actor( } }; - if !pending_replies.contains(&id) { + if !pending_replies.mark_published(&id) { continue; } - if let Err(error) = transport_tx.unbounded_send(TransportFrame::Single(request)) { + if let Err(error) = + publish(&transport_tx, TransportFrame::Single(request), permit).await + { let error = crate::Error::into_internal_error(error); if let Some(pending_reply) = pending_replies.remove(&id) { pending_reply.fail(error.clone()); @@ -147,32 +241,14 @@ pub(super) async fn outgoing_protocol_actor( continue; } OutgoingMessage::Notification { untyped } => { - let messages = match protocol_compat.outgoing_notification(untyped) { - Ok(messages) => messages, - Err(error) => { - tracing::warn!( - ?error, - "Dropping outgoing notification after preparation failed" - ); - continue; - } - }; - - for untyped in messages { - let message = match untyped.into_raw_jsonrpc_message(None) { - Ok(message) => message, - Err(error) => { - tracing::warn!( - ?error, - "Dropping outgoing notification after serialization failed" - ); - continue; - } - }; - transport_tx - .unbounded_send(TransportFrame::Single(message)) - .map_err(crate::Error::into_internal_error)?; - } + publish_notification( + &transport_tx, + &protocol_compat, + &pending_replies, + untyped, + permit, + ) + .await?; continue; } OutgoingMessage::Response { @@ -198,12 +274,13 @@ pub(super) async fn outgoing_protocol_actor( destination, ) } + OutgoingMessage::Admitted { .. } => { + unreachable!("application admission is unwrapped above") + } }; - if let Some(frame) = destination.complete(json_rpc_message) { - transport_tx - .unbounded_send(frame) - .map_err(crate::Error::into_internal_error)?; + if let Some((frame, permit)) = destination.complete_admitted(json_rpc_message, permit) { + publish(&transport_tx, frame, permit).await?; } } diff --git a/src/agent-client-protocol/src/jsonrpc/raw_error.rs b/src/agent-client-protocol/src/jsonrpc/raw_error.rs new file mode 100644 index 00000000..00a6ec97 --- /dev/null +++ b/src/agent-client-protocol/src/jsonrpc/raw_error.rs @@ -0,0 +1,177 @@ +//! Transport-level error objects, before choosing an application protocol. + +use agent_client_protocol_schema::MaybeUndefined; +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; + +/// A JSON-RPC error without ACP-specific interpretation. +/// +/// Raw transports and relays preserve unknown fields and distinguish omitted +/// `data` from explicit JSON null. Convert to [`crate::Error`] only when +/// dispatching an ACP response; an MCP error code belongs to a different domain. +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] +#[non_exhaustive] +pub struct RawJsonRpcError { + /// The peer's numeric error code, not an ACP [`crate::ErrorCode`]. + pub code: i32, + /// The peer's error message. + pub message: String, + /// Optional error data. Explicit null is retained separately from omission. + #[serde(default, skip_serializing_if = "MaybeUndefined::is_undefined")] + pub data: MaybeUndefined, + /// Additional fields on the error object. + #[serde(flatten)] + pub extra: Map, +} + +/// A transport-level JSON-RPC response with an opaque result or raw error. +/// +/// Errors are boxed so their extensible representation does not enlarge every +/// request, notification, and queued frame. +pub type RawJsonRpcResponse = + agent_client_protocol_schema::rpc::Response>; + +impl RawJsonRpcError { + /// Construct an error without data or extension fields. + #[must_use] + pub fn new(code: i32, message: impl Into) -> Self { + Self { + code, + message: message.into(), + data: MaybeUndefined::Undefined, + extra: Map::new(), + } + } + + /// Set error data, preserving explicit null. + #[must_use] + pub fn data(mut self, data: Value) -> Self { + self.data = if data.is_null() { + MaybeUndefined::Null + } else { + MaybeUndefined::Value(data) + }; + self + } + + /// Interpret this error as an ACP response for the typed dispatcher. + /// + /// ACP's error type does not model extension fields, so this intentionally + /// discards `extra`. Do not use it when forwarding raw frames or projecting + /// errors from another protocol such as MCP. + #[must_use] + pub fn into_acp_error(self) -> crate::Error { + let mut error = crate::Error::new(self.code, self.message); + error.data = match self.data { + MaybeUndefined::Undefined => None, + MaybeUndefined::Null => Some(Value::Null), + MaybeUndefined::Value(data) => Some(data), + }; + error + } +} + +impl From for RawJsonRpcError { + fn from(error: crate::Error) -> Self { + let raw = Self::new(error.code.into(), error.message); + match error.data { + Some(data) => raw.data(data), + None => raw, + } + } +} + +impl std::fmt::Display for RawJsonRpcError { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(formatter, "{} ({})", self.message, self.code) + } +} + +impl std::error::Error for RawJsonRpcError {} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{Channel, RawJsonRpcMessage, TransportFrame}; + use futures::{SinkExt as _, StreamExt as _}; + use serde_json::json; + + #[test] + fn raw_errors_preserve_omission_null_values_and_extensions() { + for data in [None, Some(Value::Null), Some(json!({"detail":[1,2]}))] { + let mut error = json!({ + "code": -32000, + "message": "peer", + "extension": {"retry": true}, + "_meta": {"opaque": "kept"} + }); + if let Some(data) = &data { + error["data"] = data.clone(); + } + let wire = json!({"jsonrpc":"2.0", "id":"logical", "error":error}); + let parsed: RawJsonRpcMessage = serde_json::from_value(wire.clone()).unwrap(); + assert_eq!(serde_json::to_value(&parsed).unwrap(), wire); + let RawJsonRpcMessage::Response(RawJsonRpcResponse::Error { error, .. }) = parsed + else { + panic!("expected a raw error response"); + }; + assert_eq!(error.code, -32000); + match data { + None => assert!(error.data.is_undefined()), + Some(Value::Null) => assert!(error.data.is_null()), + Some(value) => assert_eq!(error.data.value(), Some(&value)), + } + } + } + + #[tokio::test] + async fn raw_error_batch_survives_framing_and_budgeted_relay() { + let wire = json!([ + {"jsonrpc":"2.0", "id":"error", "error":{ + "code":-32000, "message":"peer", "data":null, "extension":{"retry":true} + }}, + {"jsonrpc":"2.0", "id":"success", "result":null} + ]); + let frame = TransportFrame::parse_json(&wire.to_string()); + assert!(matches!(&frame, TransportFrame::Batch(_))); + let (source, mut relay_in) = Channel::duplex(); + let (mut relay_out, mut destination) = Channel::duplex(); + source.tx.send_frame(frame).await.unwrap(); + let admitted = relay_in.rx.next().await.unwrap(); + relay_out.tx.send(admitted).await.unwrap(); + let received = destination.rx.next().await.unwrap(); + let received: Value = serde_json::from_str(&received.frame().to_json().unwrap()).unwrap(); + assert_eq!(received, wire); + } + + #[test] + fn acp_error_interpretation_is_explicit_and_keeps_data_presence() { + let raw = RawJsonRpcError::new(-32000, "peer"); + assert_eq!(raw.clone().into_acp_error().data, None); + let error = raw.data(Value::Null).into_acp_error(); + assert_eq!(error.code, crate::ErrorCode::AuthRequired); + assert_eq!(error.data, Some(Value::Null)); + let roundtrip = RawJsonRpcError::from(error); + assert!(roundtrip.data.is_null()); + assert!(roundtrip.extra.is_empty()); + } + + #[test] + fn malformed_raw_errors_are_still_rejected() { + for error in [ + Value::Null, + json!({"code":-32000}), + json!({"message":"peer"}), + json!({"code":null, "message":"peer"}), + json!({"code":1.5, "message":"peer"}), + json!({"code":-32000, "message":null}), + ] { + assert!( + serde_json::from_value::( + json!({"jsonrpc":"2.0", "id":1, "error":error}) + ) + .is_err() + ); + } + } +} diff --git a/src/agent-client-protocol/src/jsonrpc/run.rs b/src/agent-client-protocol/src/jsonrpc/run.rs index 571bf61a..f5aaa870 100644 --- a/src/agent-client-protocol/src/jsonrpc/run.rs +++ b/src/agent-client-protocol/src/jsonrpc/run.rs @@ -7,6 +7,8 @@ use std::future::Future; use std::marker::PhantomData; +use futures::future::{Either, select}; + use crate::{ ConnectionTo, jsonrpc::{ConnectionContext, RawConnectionContext, connection_context}, @@ -67,8 +69,29 @@ where // Box the futures to avoid stack overflow with deeply nested RunIn chains let a_fut = Box::pin(self.a.run_with_connection_to(cx.clone())); let b_fut = Box::pin(self.b.run_with_connection_to(cx.clone())); - let ((), ()) = futures::future::try_join(a_fut, b_fut).await?; - Ok(()) + match select(a_fut, b_fut).await { + Either::Left((Ok(()), b)) => b.await, + Either::Right((Ok(()), a)) => a.await, + Either::Left((Err(error), b)) => { + cx.request_shutdown(); + // A different runner may own cleanup of a scoped MCP tool. + // Continue polling it without waiting for an unrelated + // never-ending runner after protected operations finish. + match select(b, Box::pin(cx.wait_protected_operations())).await { + Either::Left((_, cleanup)) => cleanup.await, + Either::Right(((), _)) => {} + } + Err(error) + } + Either::Right((Err(error), a)) => { + cx.request_shutdown(); + match select(a, Box::pin(cx.wait_protected_operations())).await { + Either::Left((_, cleanup)) => cleanup.await, + Either::Right(((), _)) => {} + } + Err(error) + } + } } } diff --git a/src/agent-client-protocol/src/jsonrpc/task_actor.rs b/src/agent-client-protocol/src/jsonrpc/task_actor.rs index 92a05a6d..8f4e8289 100644 --- a/src/agent-client-protocol/src/jsonrpc/task_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/task_actor.rs @@ -1,16 +1,48 @@ use std::panic::Location; +use std::sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, +}; -use futures::{FutureExt, channel::mpsc, future::BoxFuture}; +use futures::channel::oneshot; +use futures::future::{self, Either}; +use futures::{FutureExt, StreamExt, future::BoxFuture}; use crate::ConnectionTo; use crate::role::Role; -use crate::util::process_stream_concurrently; -pub type TaskTx = mpsc::UnboundedSender; +#[derive(Clone, Debug)] +pub struct TaskTx { + sender: super::admission::Sender, + live: Arc, + capacity: usize, +} + +pub fn task_channel(capacity: usize) -> (TaskTx, super::admission::SimpleReceiver) { + let capacity = capacity.max(1); + let (sender, receiver) = super::admission::channel_with_capacity(capacity); + ( + TaskTx { + sender, + live: Arc::new(AtomicUsize::new(0)), + capacity, + }, + receiver, + ) +} + +struct LiveTask(Arc); + +impl Drop for LiveTask { + fn drop(&mut self) { + self.0.fetch_sub(1, Ordering::AcqRel); + } +} #[must_use] pub(crate) struct Task { future: BoxFuture<'static, Result<(), crate::Error>>, + live: Option, } impl Task { @@ -35,12 +67,21 @@ impl Task { } }, ) - .boxed() + .boxed(), + live: None, } } - pub fn spawn(self, task_tx: &TaskTx) -> Result<(), crate::Error> { + pub fn spawn(mut self, task_tx: &TaskTx) -> Result<(), crate::Error> { + task_tx + .live + .fetch_update(Ordering::AcqRel, Ordering::Acquire, |live| { + (live < task_tx.capacity).then_some(live + 1) + }) + .map_err(|_| crate::util::internal_error("live task capacity exceeded"))?; + self.live = Some(LiveTask(task_tx.live.clone())); task_tx + .sender .unbounded_send(self) .map_err(crate::util::internal_error)?; Ok(()) @@ -54,13 +95,80 @@ impl Task { /// The "task actor" manages dynamically spawned tasks. pub(super) async fn task_actor( - task_rx: mpsc::UnboundedReceiver, - _cx: &ConnectionTo, + task_rx: super::admission::SimpleReceiver, + cx: &ConnectionTo, + max_running_tasks: usize, ) -> Result<(), crate::Error> { - process_stream_concurrently( - task_rx, - async |task| task.future.await, - |a, b| Box::pin(a(b)), - ) - .await + let (error_tx, error_rx) = oneshot::channel(); + let first_error = Arc::new(Mutex::new(Some(error_tx))); + let running = task_rx.for_each_concurrent(max_running_tasks.max(1), |task| { + let first_error = first_error.clone(); + async move { + let Task { future, live } = task; + let result = future.await; + drop(live); + if let Err(error) = result + && let Some(tx) = first_error + .lock() + .expect("task error mutex poisoned") + .take() + { + drop(tx.send(error)); + } + } + }); + let on_error = async { + let error = error_rx + .await + .expect("task driver dropped before completion"); + cx.incoming_closed.request_shutdown(); + // Keep polling the driver while native supervisors finish. A failed + // disposable task cannot drop those supervisors or force us to join + // arbitrary never-ending disposable tasks. + cx.wait_protected_operations().await; + Err(error) + }; + match future::select(Box::pin(running), Box::pin(on_error)).await { + Either::Left(((), _)) => Ok(()), + Either::Right((result, _)) => result, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use futures::channel::oneshot; + + #[test] + fn running_child_occupies_total_live_capacity_not_just_queue_capacity() { + futures::executor::block_on(async { + let (tx, mut rx) = task_channel(1); + let (child_done_tx, child_done_rx) = oneshot::channel::<()>(); + Task::new(Location::caller(), async move { + let _ = child_done_rx.await; + Ok(()) + }) + .spawn(&tx) + .unwrap(); + // The child is no longer in the waiting queue, but remains live. + let child = rx.next().await.unwrap(); + let (callback_dropped_tx, callback_dropped_rx) = oneshot::channel::<()>(); + let callback = async move { + let _drop_on_rejection = callback_dropped_tx; + futures::future::pending::<()>().await; + Ok(()) + }; + let rejection = Task::new(Location::caller(), callback) + .spawn(&tx) + .expect_err("an ordered callback cannot wait behind a permanent child"); + assert!(rejection.to_string().contains("live task capacity")); + assert!(callback_dropped_rx.now_or_never().unwrap().is_err()); + + drop(child_done_tx); + child.run_for_test().await.unwrap(); + Task::new(Location::caller(), async { Ok(()) }) + .spawn(&tx) + .expect("child completion releases total live capacity"); + }); + } } diff --git a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs index 8737ca81..d71d4004 100644 --- a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs @@ -1,12 +1,16 @@ use std::pin::pin; // Types re-exported from crate root +use crate::RawJsonRpcResponse as Response; use crate::jsonrpc::{RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame}; -use crate::schema::v1::Response; use futures::StreamExt as _; -use futures::channel::mpsc; use serde::Deserialize as _; +/// Maximum bytes in one wire value (excluding its newline). +pub const MAX_FRAME_BYTES: usize = 16 * 1024 * 1024; +/// Maximum number of JSON-RPC values carried in one batch. +pub const MAX_BATCH_ENTRIES: usize = 64; + enum ParsedIncomingLine { Single(RawJsonRpcMessage), Malformed { raw: String, error: crate::Error }, @@ -14,13 +18,21 @@ enum ParsedIncomingLine { } fn parse_incoming_line(line: &str) -> ParsedIncomingLine { + if line.len() > MAX_FRAME_BYTES { + return ParsedIncomingLine::Malformed { + raw: String::new(), + error: crate::Error::invalid_request().data("JSON-RPC frame exceeds maximum size"), + }; + } let value = match serde_json::from_str::(line) { Ok(value) => value, Err(error) => { tracing::debug!(?error, "Failed to parse incoming JSON-RPC JSON"); return ParsedIncomingLine::Malformed { raw: line.to_owned(), - error: crate::Error::parse_error().data(serde_json::json!({ "line": line })), + error: crate::Error::parse_error().data(serde_json::json!({ + "line": line.chars().take(256).collect::() + })), }; } }; @@ -30,6 +42,12 @@ fn parse_incoming_line(line: &str) -> ParsedIncomingLine { raw: line.to_owned(), error: crate::Error::invalid_request(), }, + serde_json::Value::Array(entries) if entries.len() > MAX_BATCH_ENTRIES => { + ParsedIncomingLine::Malformed { + raw: String::new(), + error: crate::Error::invalid_request().data("JSON-RPC batch exceeds maximum width"), + } + } serde_json::Value::Array(entries) => { let entries = entries .into_iter() @@ -60,6 +78,62 @@ fn parse_incoming_line(line: &str) -> ParsedIncomingLine { } } +/// Read newline-delimited UTF-8 without allocating an unterminated line larger +/// than the frame budget. An oversized line terminates the transport explicitly. +pub fn bounded_lines( + input: R, +) -> impl futures::Stream> { + use futures::io::BufReader; + use futures::{AsyncBufReadExt, stream}; + stream::unfold(Some(BufReader::new(input)), |reader| async move { + let mut reader = reader?; + let mut bytes = Vec::new(); + loop { + let chunk = match reader.fill_buf().await { + Ok(chunk) => chunk, + Err(error) => return Some((Err(error), None)), + }; + if chunk.is_empty() { + return if bytes.is_empty() { + None + } else { + Some(( + String::from_utf8(bytes) + .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e)), + None, + )) + }; + } + let width = chunk + .iter() + .position(|&b| b == b'\n') + .map_or(chunk.len(), |i| i + 1); + if bytes.len() + width > MAX_FRAME_BYTES + 1 { + return Some(( + Err(std::io::Error::new( + std::io::ErrorKind::InvalidData, + "JSON-RPC line exceeds maximum frame size", + )), + None, + )); + } + bytes.extend_from_slice(&chunk[..width]); + reader.consume_unpin(width); + if bytes.last() == Some(&b'\n') { + bytes.pop(); + if bytes.last() == Some(&b'\r') { + bytes.pop(); + } + return Some(( + String::from_utf8(bytes) + .map_err(|e| std::io::Error::new(std::io::ErrorKind::InvalidData, e)), + Some(reader), + )); + } + } + }) +} + impl TransportFrame { /// Parse one JSON-RPC wire value while preserving batch boundaries. /// @@ -110,19 +184,21 @@ impl TransportFrame { /// /// This is the transport layer - it has no knowledge of protocol semantics (IDs, correlation, etc.). async fn transport_outgoing_frames_actor( - transport_rx: impl futures::Stream, + transport_rx: impl futures::Stream, outgoing_lines: impl futures::Sink, ) -> Result<(), crate::Error> { use futures::SinkExt; let mut transport_rx = pin!(transport_rx); let mut outgoing_lines = pin!(outgoing_lines); - while let Some(frame) = transport_rx.next().await { + while let Some(budgeted) = transport_rx.next().await { + let (frame, _permit) = budgeted.into_parts(); let json_rpc_message = match frame { TransportFrame::Single(message) => message, TransportFrame::Malformed { raw, .. } => { let raw = malformed_line_value(raw)?; tracing::trace!(message = ?raw, "Relaying invalid JSON-RPC value"); + ensure_frame_size(&raw)?; outgoing_lines .send(raw) .await @@ -133,6 +209,7 @@ async fn transport_outgoing_frames_actor( let line = serde_json::to_string(&batch).map_err(crate::Error::into_internal_error)?; tracing::trace!(message = %line, "Sending JSON-RPC batch"); + ensure_frame_size(&line)?; outgoing_lines .send(line) .await @@ -143,6 +220,7 @@ async fn transport_outgoing_frames_actor( match serde_json::to_string(&json_rpc_message) { Ok(line) => { tracing::trace!(message = %line, "Sending JSON-RPC message"); + ensure_frame_size(&line)?; outgoing_lines .send(line) .await @@ -177,6 +255,7 @@ async fn transport_outgoing_frames_actor( Err(crate::Error::internal_error()), )) .unwrap(); + ensure_frame_size(&error_line)?; outgoing_lines .send(error_line) .await @@ -189,6 +268,14 @@ async fn transport_outgoing_frames_actor( Ok(()) } +fn ensure_frame_size(line: &str) -> Result<(), crate::Error> { + if line.len() > MAX_FRAME_BYTES { + Err(crate::Error::invalid_request().data("outgoing JSON-RPC frame exceeds maximum size")) + } else { + Ok(()) + } +} + fn malformed_line_value(raw: String) -> Result { if !raw.contains('\r') && !raw.contains('\n') { return Ok(raw); @@ -202,7 +289,7 @@ fn malformed_line_value(raw: String) -> Result { } pub(super) async fn transport_outgoing_lines_actor( - transport_rx: mpsc::UnboundedReceiver, + transport_rx: super::FrameReceiver, outgoing_lines: impl futures::Sink, ) -> Result<(), crate::Error> { transport_outgoing_frames_actor(transport_rx, outgoing_lines).await @@ -222,7 +309,7 @@ pub(super) async fn transport_outgoing_lines_actor( /// This is the transport layer - it has no knowledge of protocol semantics. pub(super) async fn transport_incoming_lines_actor( incoming_lines: impl futures::Stream>, - transport_tx: mpsc::UnboundedSender, + transport_tx: super::FrameSender, ) -> Result<(), crate::Error> { let mut incoming_lines = pin!(incoming_lines); while let Some(line_result) = incoming_lines.next().await { @@ -232,17 +319,20 @@ pub(super) async fn transport_incoming_lines_actor( match parse_incoming_line(&line) { ParsedIncomingLine::Single(message) => { transport_tx - .unbounded_send(TransportFrame::Single(message)) + .send_frame(TransportFrame::Single(message)) + .await .map_err(crate::Error::into_internal_error)?; } ParsedIncomingLine::Malformed { raw, error } => { transport_tx - .unbounded_send(TransportFrame::Malformed { raw, error }) + .send_frame(TransportFrame::Malformed { raw, error }) + .await .map_err(crate::Error::into_internal_error)?; } ParsedIncomingLine::Batch(entries) => { transport_tx - .unbounded_send(TransportFrame::Batch(entries)) + .send_frame(TransportFrame::Batch(entries)) + .await .map_err(crate::Error::into_internal_error)?; } } @@ -257,6 +347,28 @@ mod tests { use super::*; use crate::ErrorCode; + #[test] + fn rejects_batches_over_width_limit() { + let batch = format!("[{}]", vec!["null"; MAX_BATCH_ENTRIES + 1].join(",")); + let ParsedIncomingLine::Malformed { error, .. } = parse_incoming_line(&batch) else { + panic!("oversized batch must be rejected"); + }; + assert_eq!(error.code, ErrorCode::InvalidRequest); + } + + #[tokio::test] + async fn oversized_unterminated_line_fails_before_eof() { + let input = futures::io::Cursor::new(vec![b'x'; MAX_FRAME_BYTES + 2]); + let mut lines = Box::pin(bounded_lines(input)); + let error = lines + .next() + .await + .expect("explicit framing failure") + .unwrap_err(); + assert_eq!(error.kind(), std::io::ErrorKind::InvalidData); + assert!(lines.next().await.is_none()); + } + #[test] fn parses_batch_entries_independently() { let ParsedIncomingLine::Batch(batch) = parse_incoming_line( @@ -446,15 +558,19 @@ mod tests { Ok::<_, std::io::Error>(captured) }); - transport_outgoing_frames_actor( - futures::stream::iter([TransportFrame::Malformed { + let (source, destination) = crate::Channel::duplex(); + source + .tx + .send_frame(TransportFrame::Malformed { raw: raw.clone(), error: crate::Error::parse_error(), - }]), - outgoing, - ) - .await - .unwrap(); + }) + .await + .unwrap(); + drop(source); + transport_outgoing_frames_actor(destination.rx, outgoing) + .await + .unwrap(); let lines = captured.lock().unwrap(); assert_eq!(lines.len(), 1); diff --git a/src/agent-client-protocol/src/lib.rs b/src/agent-client-protocol/src/lib.rs index 0a543aed..7f58393e 100644 --- a/src/agent-client-protocol/src/lib.rs +++ b/src/agent-client-protocol/src/lib.rs @@ -143,10 +143,12 @@ pub mod util; pub use capabilities::*; pub use jsonrpc::{ - Builder, ByteStreams, Channel, ConnectionContext, ConnectionTo, Dispatch, DynamicHandlerGuard, - HandleConnectionClose, HandleDispatchFrom, Handled, INCOMING_TRANSPORT_CLOSED_REASON, - IntoHandled, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, Lines, - NullClose, NullHandler, RawConnectionContext, RawJsonRpcMessage, RawJsonRpcParams, Responder, + BudgetedFrame, Builder, ByteStreams, Channel, ConnectionContext, ConnectionLimits, + ConnectionTo, Dispatch, DynamicHandlerGuard, FrameAdmission, FramePermit, FrameReceiver, + FrameSender, HandleConnectionClose, HandleDispatchFrom, Handled, + INCOMING_TRANSPORT_CLOSED_REASON, IntoHandled, JsonRpcMessage, JsonRpcNotification, + JsonRpcRequest, JsonRpcResponse, Lines, NullClose, NullHandler, RawConnectionContext, + RawJsonRpcError, RawJsonRpcMessage, RawJsonRpcParams, RawJsonRpcResponse, Responder, ResponseRouter, SentRequest, TransportBatch, TransportBatchEntry, TransportFrame, UntypedMessage, is_incoming_transport_closed, run::{ChainRun, NullRun, RunWithConnectionTo}, @@ -162,7 +164,7 @@ pub use role::{ acp::{Agent, Client, Conductor, Proxy}, }; -pub use component::{ConnectTo, DynConnectTo}; +pub use component::{ConnectTo, ConnectionDriver, DynConnectTo}; /// Implementation details used by the derive macros. #[doc(hidden)] diff --git a/src/agent-client-protocol/src/mcp_server/active_session.rs b/src/agent-client-protocol/src/mcp_server/active_session.rs index c30322cf..28d90f3e 100644 --- a/src/agent-client-protocol/src/mcp_server/active_session.rs +++ b/src/agent-client-protocol/src/mcp_server/active_session.rs @@ -1,213 +1,248 @@ -use std::{marker::PhantomData, sync::Arc}; +//! Request-scoped native MCP transport. Each ACP request owns execution and cleanup. -use futures::channel::mpsc; -use futures::{SinkExt, StreamExt}; -use rustc_hash::FxHashMap; +use futures::{ + StreamExt, + channel::oneshot, + future::{self, Either}, +}; use serde_json::{Map, Value}; - -use crate::mcp_server::{McpConnectionContext, McpConnectionTo, McpServerConnect}; -use crate::role; -use crate::role::HasPeer; -use crate::schema::v1::{ - ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, DisconnectMcpResponse, - McpConnectionId, McpServerAcpId, MessageMcpNotification, MessageMcpRequest, MessageMcpResponse, +use std::{ + collections::HashMap, + io::Write, + marker::PhantomData, + sync::{Arc, Mutex, Weak}, }; -use crate::util::MatchDispatchFrom; + use crate::{ Agent, Channel, ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, - JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, Responder, Role, UntypedMessage, + JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, RawJsonRpcError, RawJsonRpcMessage, + RawJsonRpcResponse, Responder, Role, TransportFrame, + mcp_server::{ + MCP_BACKEND_FAILURE, MCP_RESOURCE_EXHAUSTED, McpConnectionContext, McpConnectionTo, + McpOperationCancellation, McpOutcome, McpRequest, McpRequestContext, McpServerConnect, + McpService, + }, + role::HasPeer, + schema::v1::{ + McpError, McpRequestId, McpServerAcpId, MessageMcpNotification, MessageMcpRequest, + MessageMcpResponse, RequestId, + }, + util::MatchDispatchFrom, }; -/// Stable protocol v1 native MCP-over-ACP wire types. -pub(super) struct V1McpProtocol; +// These bound admitted work and individual payloads. +const MAX_ACTIVE_REQUESTS: usize = 64; +const MAX_PAYLOAD_BYTES: usize = 16 * 1024 * 1024; +const MCP_VERSION: &str = "2026-07-28"; +type ActiveRequests = Arc>>>; -/// Draft protocol v2 native MCP-over-ACP wire types. +pub(super) struct V1McpProtocol; #[cfg(feature = "unstable_protocol_v2")] pub(super) struct V2McpProtocol; -pub(super) struct McpMessage { - method: String, - params: Option>, -} - pub(super) trait McpProtocol: Send + 'static { - type ConnectRequest: JsonRpcRequest; - type ConnectResponse: JsonRpcResponse; type MessageRequest: JsonRpcRequest; + type MessageResponse: JsonRpcResponse + serde::Serialize; type MessageNotification: JsonRpcNotification; - type MessageResponse: JsonRpcResponse; - type DisconnectRequest: JsonRpcRequest; - type DisconnectResponse: JsonRpcResponse; - - fn connect_server_id(request: &Self::ConnectRequest) -> McpServerAcpId; - fn connect_response(connection_id: McpConnectionId) -> Self::ConnectResponse; - fn message_request( - connection_id: McpConnectionId, - method: String, - params: Option>, - ) -> Self::MessageRequest; - fn message_notification( - connection_id: McpConnectionId, + + fn response(outcome: McpOutcome) -> Self::MessageResponse; + fn server_id(request: &Self::MessageRequest) -> McpServerAcpId; + fn request_id(request: &Self::MessageRequest) -> McpRequestId; + fn into_request(request: Self::MessageRequest) -> (String, Option>); + fn notification( + server_id: McpServerAcpId, + request_id: McpRequestId, method: String, params: Option>, ) -> Self::MessageNotification; - fn message_request_connection_id(request: &Self::MessageRequest) -> McpConnectionId; - fn message_notification_connection_id( - notification: &Self::MessageNotification, - ) -> McpConnectionId; - fn into_message_request(request: Self::MessageRequest) -> McpMessage; - fn into_message_notification(notification: Self::MessageNotification) -> McpMessage; - fn disconnect_connection_id(request: &Self::DisconnectRequest) -> McpConnectionId; - fn disconnect_response() -> Self::DisconnectResponse; } impl McpProtocol for V1McpProtocol { - type ConnectRequest = ConnectMcpRequest; - type ConnectResponse = ConnectMcpResponse; type MessageRequest = MessageMcpRequest; - type MessageNotification = MessageMcpNotification; type MessageResponse = MessageMcpResponse; - type DisconnectRequest = DisconnectMcpRequest; - type DisconnectResponse = DisconnectMcpResponse; + type MessageNotification = MessageMcpNotification; - fn connect_server_id(request: &Self::ConnectRequest) -> McpServerAcpId { - request.server_id.clone() + fn response(outcome: McpOutcome) -> Self::MessageResponse { + match outcome { + McpOutcome::Result(value) => MessageMcpResponse::success(value), + McpOutcome::Error(error) => MessageMcpResponse::error(error), + } } - fn connect_response(connection_id: McpConnectionId) -> Self::ConnectResponse { - ConnectMcpResponse::new(connection_id) + fn server_id(request: &Self::MessageRequest) -> McpServerAcpId { + request.server_id.clone() } - - fn message_request( - connection_id: McpConnectionId, - method: String, - params: Option>, - ) -> Self::MessageRequest { - MessageMcpRequest::new(connection_id, method).params(params) + fn request_id(request: &Self::MessageRequest) -> McpRequestId { + request.request_id.clone() } - - fn message_notification( - connection_id: McpConnectionId, + fn into_request(request: Self::MessageRequest) -> (String, Option>) { + (request.method, request.params) + } + fn notification( + server_id: McpServerAcpId, + request_id: McpRequestId, method: String, params: Option>, ) -> Self::MessageNotification { - MessageMcpNotification::new(connection_id, method).params(params) + MessageMcpNotification::new(server_id, request_id, method).params(params) } +} - fn message_request_connection_id(request: &Self::MessageRequest) -> McpConnectionId { - request.connection_id.clone() - } +fn into_mcp_error(error: impl Into) -> McpError { + let error = error.into(); + let mut mcp = McpError::new(error.code, error.message); + mcp.data = error.data; + mcp.extra = error.extra; + mcp +} - fn message_notification_connection_id( - notification: &Self::MessageNotification, - ) -> McpConnectionId { - notification.connection_id.clone() - } +fn outcome_response( + outcome: McpOutcome, +) -> Result { + let response = Protocol::response(outcome); + check_payload_size(&response, MAX_PAYLOAD_BYTES)?; + Ok(response) +} - fn into_message_request(request: Self::MessageRequest) -> McpMessage { - McpMessage { - method: request.method, - params: request.params, +fn project_outcome(outcome: McpOutcome, is_discovery: bool) -> Result { + match outcome { + McpOutcome::Result(mut value) if is_discovery => { + constrain_discovery_versions(&mut value)?; + Ok(McpOutcome::Result(value)) } + other => Ok(other), } +} - fn into_message_notification(notification: Self::MessageNotification) -> McpMessage { - McpMessage { - method: notification.method, - params: notification.params, +fn send_outcome( + responder: Responder, + result: Result, + is_discovery: bool, +) -> Result<(), crate::Error> { + match result { + Ok(outcome) => { + // Projection failures are MCP outcomes; binding and size failures + // remain named outer ACP errors, regardless of backend type. + let outcome = project_outcome(outcome, is_discovery) + .unwrap_or_else(|error| McpOutcome::Error(into_mcp_error(error))); + match outcome_response::(outcome) { + Ok(response) => responder.respond(response), + Err(error) => responder.respond_with_error(error), + } } - } - - fn disconnect_connection_id(request: &Self::DisconnectRequest) -> McpConnectionId { - request.connection_id.clone() - } - - fn disconnect_response() -> Self::DisconnectResponse { - DisconnectMcpResponse::new() + Err(error) => responder.respond_with_error(error), } } #[cfg(feature = "unstable_protocol_v2")] impl McpProtocol for V2McpProtocol { - type ConnectRequest = crate::schema::v2::ConnectMcpRequest; - type ConnectResponse = crate::schema::v2::ConnectMcpResponse; type MessageRequest = crate::schema::v2::MessageMcpRequest; - type MessageNotification = crate::schema::v2::MessageMcpNotification; type MessageResponse = crate::schema::v2::MessageMcpResponse; - type DisconnectRequest = crate::schema::v2::DisconnectMcpRequest; - type DisconnectResponse = crate::schema::v2::DisconnectMcpResponse; + type MessageNotification = crate::schema::v2::MessageMcpNotification; - fn connect_server_id(request: &Self::ConnectRequest) -> McpServerAcpId { - McpServerAcpId::new(request.server_id.0.clone()) + fn response(outcome: McpOutcome) -> Self::MessageResponse { + match outcome { + McpOutcome::Result(value) => Self::MessageResponse::success(value), + McpOutcome::Error(error) => { + // The service outcome uses the v1 error representation. Adapt it + // explicitly here instead of coupling the versioned wire types. + let mut wire_error = crate::schema::v2::McpError::new(error.code, error.message); + wire_error.data = error.data; + wire_error.extra = error.extra; + Self::MessageResponse::error(wire_error) + } + } } - fn connect_response(connection_id: McpConnectionId) -> Self::ConnectResponse { - crate::schema::v2::ConnectMcpResponse::new(connection_id.0) + fn server_id(request: &Self::MessageRequest) -> McpServerAcpId { + McpServerAcpId::new(request.server_id.0.clone()) } - - fn message_request( - connection_id: McpConnectionId, - method: String, - params: Option>, - ) -> Self::MessageRequest { - crate::schema::v2::MessageMcpRequest::new(connection_id.0, method).params(params) + fn request_id(request: &Self::MessageRequest) -> McpRequestId { + McpRequestId::new(request.request_id.0.clone()) } - - fn message_notification( - connection_id: McpConnectionId, + fn into_request(request: Self::MessageRequest) -> (String, Option>) { + (request.method, request.params) + } + fn notification( + server_id: McpServerAcpId, + request_id: McpRequestId, method: String, params: Option>, ) -> Self::MessageNotification { - crate::schema::v2::MessageMcpNotification::new(connection_id.0, method).params(params) - } - - fn message_request_connection_id(request: &Self::MessageRequest) -> McpConnectionId { - McpConnectionId::new(request.connection_id.0.clone()) + crate::schema::v2::MessageMcpNotification::new(server_id.0, request_id.0, method) + .params(params) } +} - fn message_notification_connection_id( - notification: &Self::MessageNotification, - ) -> McpConnectionId { - McpConnectionId::new(notification.connection_id.0.clone()) - } +/// Active operations belong to the handler; dropping the declaration closes every operation. +pub(super) struct McpActiveSession { + server_id: McpServerAcpId, + mcp_connect: Arc>, + service: Option>>, + active: ActiveRequests, + protocol: PhantomData Protocol>, +} - fn into_message_request(request: Self::MessageRequest) -> McpMessage { - McpMessage { - method: request.method, - params: request.params, - } - } +struct ActiveRequest { + active: Weak>>>, + id: McpRequestId, +} - fn into_message_notification(notification: Self::MessageNotification) -> McpMessage { - McpMessage { - method: notification.method, - params: notification.params, +impl Drop for ActiveRequest { + fn drop(&mut self) { + if let Some(active) = self.active.upgrade() { + active + .lock() + .expect("MCP request registry poisoned") + .remove(&self.id); } } +} - fn disconnect_connection_id(request: &Self::DisconnectRequest) -> McpConnectionId { - McpConnectionId::new(request.connection_id.0.clone()) +fn admit_request( + active: &ActiveRequests, + id: McpRequestId, +) -> Result<(ActiveRequest, oneshot::Receiver<()>), crate::Error> { + let (stop_tx, stop_rx) = oneshot::channel(); + let mut requests = active.lock().expect("MCP request registry poisoned"); + if requests.contains_key(&id) { + return Err(crate::Error::invalid_params().data("duplicate active MCP requestId")); } - - fn disconnect_response() -> Self::DisconnectResponse { - crate::schema::v2::DisconnectMcpResponse::new() + if requests.len() >= MAX_ACTIVE_REQUESTS { + return Err( + crate::Error::new(MCP_RESOURCE_EXHAUSTED, "MCP active request limit exceeded") + .data(serde_json::json!({"limit": MAX_ACTIVE_REQUESTS})), + ); } + requests.insert(id.clone(), stop_tx); + Ok(( + ActiveRequest { + active: Arc::downgrade(active), + id, + }, + stop_rx, + )) } -/// The message handler for an MCP server offered to a particular session. -/// This is added as a dynamic handler to the connection context and handles -/// native MCP-over-ACP messages for the declared server ID. -pub(super) struct McpActiveSession { - /// The opaque ACP transport identifier for this MCP server. - server_id: McpServerAcpId, - - /// The MCP server we are managing. - mcp_connect: Arc>, - - /// Active connections to MCP server tasks. - connections: FxHashMap>, - - protocol: PhantomData Protocol>, +/// Count serialized bytes without allocating another copy of a potentially large payload. +fn check_payload_size(value: &impl serde::Serialize, limit: usize) -> Result<(), crate::Error> { + struct Budget(usize); + impl Write for Budget { + fn write(&mut self, bytes: &[u8]) -> std::io::Result { + self.0 = self + .0 + .checked_sub(bytes.len()) + .ok_or_else(|| std::io::Error::other("MCP payload limit exceeded"))?; + Ok(bytes.len()) + } + fn flush(&mut self) -> std::io::Result<()> { + Ok(()) + } + } + serde_json::to_writer(Budget(limit), value).map_err(|_| { + crate::Error::new(MCP_RESOURCE_EXHAUSTED, "MCP payload limit exceeded") + .data(serde_json::json!({"limitBytes": limit})) + }) } impl McpActiveSession @@ -215,233 +250,333 @@ where Counterpart: HasPeer, Protocol: McpProtocol, { - pub fn new( + pub fn new_with_service( server_id: McpServerAcpId, mcp_connect: Arc>, + service: Option>>, ) -> Self { Self { server_id, mcp_connect, - connections: FxHashMap::default(), + service, + active: Arc::default(), protocol: PhantomData, } } - /// Handle a connection request for our MCP server by creating a new MCP connection. - fn handle_connect_request( + fn handle_request( &mut self, - request: Protocol::ConnectRequest, - responder: Responder, - acp_connection: &ConnectionTo, + request: Protocol::MessageRequest, + responder: Responder, + connection: &ConnectionTo, ) -> Result< Handled<( - Protocol::ConnectRequest, - Responder, + Protocol::MessageRequest, + Responder, )>, crate::Error, > { - let server_id = Protocol::connect_server_id(&request); + let server_id = Protocol::server_id(&request); if server_id != self.server_id { return Ok(Handled::No { message: (request, responder), retry: false, }); } - - let connection_id = - McpConnectionId::new(format!("mcp-over-acp-connection:{}", uuid::Uuid::new_v4())); - let (mcp_server_tx, mut mcp_server_rx) = mpsc::channel(128); - self.connections - .insert(connection_id.clone(), mcp_server_tx); - - let (client_channel, server_channel) = Channel::duplex(); - - let client_component = { - let connection_id = connection_id.clone(); - let acp_connection = acp_connection.clone(); - - role::mcp::Client - .builder() - .on_receive_dispatch( - async move |message: Dispatch, _mcp_connection| match message { - Dispatch::Request(request, responder) => { - let (method, params) = request.into_parts(); - let params = match into_native_params(params) { - Ok(params) => params, - Err(error) => return responder.respond_with_error(error), - }; - let request = - Protocol::message_request(connection_id.clone(), method, params); - let responder = responder.wrap_params(|method, result| { - result.and_then(|response: Protocol::MessageResponse| { - response.into_json(method) - }) - }); - let message: Dispatch< - Protocol::MessageRequest, - Protocol::MessageNotification, - > = Dispatch::Request(request, responder); - acp_connection.send_proxied_message_to(Agent, message) - } - Dispatch::Notification(notification) => { - let (method, params) = notification.into_parts(); - let params = match into_native_params(params) { - Ok(params) => params, - Err(error) => { - tracing::warn!( - ?error, - "ignoring MCP notification with positional parameters" - ); - return Ok(()); - } - }; - let notification = Protocol::message_notification( - connection_id.clone(), - method, - params, - ); - let message: Dispatch< - Protocol::MessageRequest, - Protocol::MessageNotification, - > = Dispatch::Notification(notification); - acp_connection.send_proxied_message_to(Agent, message) - } - Dispatch::Response(result, router) => router.route_with_result(result), - }, - crate::on_receive_dispatch!(), - ) - .with_spawned(move |mcp_connection| async move { - // These messages were sent by the ACP agent. Forward them to the MCP server. - while let Some(message) = mcp_server_rx.next().await { - mcp_connection.send_proxied_message_to(role::mcp::Server, message)?; - } - Ok(()) - }) - }; - - let spawned_server = self.mcp_connect.connect(McpConnectionTo { - context: McpConnectionContext::Acp { - server_id, - connection_id: connection_id.clone(), - }, - 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 }) - }); - - match spawn_results { - Ok(()) => { - responder.respond(Protocol::connect_response(connection_id))?; - Ok(Handled::Yes) - } + let request_id = Protocol::request_id(&request); + let (method, params) = Protocol::into_request(request); + if let Err(error) = check_payload_size(&(&method, ¶ms, &request_id), MAX_PAYLOAD_BYTES) + { + responder.respond_with_error(error)?; + return Ok(Handled::Yes); + } + if let Err(error) = validate_modern_request(&method, params.as_ref()) { + responder.respond(outcome_response::(McpOutcome::Error( + into_mcp_error(error), + ))?)?; + return Ok(Handled::Yes); + } + let (guard, stop_rx) = match admit_request(&self.active, request_id.clone()) { + Ok(admitted) => admitted, Err(error) => { - self.connections.remove(&connection_id); responder.respond_with_error(error)?; - Ok(Handled::Yes) + return Ok(Handled::Yes); } - } - } - - /// Forward a native MCP-over-ACP request to its MCP connection. - async fn handle_mcp_over_acp_request( - &mut self, - request: Protocol::MessageRequest, - responder: Responder, - ) -> Result< - Handled<( - Protocol::MessageRequest, - Responder, - )>, - crate::Error, - > { - let connection_id = Protocol::message_request_connection_id(&request); - let Some(mcp_server_tx) = self.connections.get_mut(&connection_id) else { - return Ok(Handled::No { - message: (request, responder), - retry: false, - }); - }; - let message = Protocol::into_message_request(request); - - let untyped = UntypedMessage { - method: message.method, - params: native_params_into_value(message.params), }; - let responder = responder.wrap_params(|method, result| { - result - .and_then(|response: Value| Protocol::MessageResponse::from_value(method, response)) - }); - mcp_server_tx - .send(Dispatch::Request(untyped, responder)) - .await - .map_err(crate::Error::into_internal_error)?; - - Ok(Handled::Yes) - } - - /// Forward a native MCP-over-ACP notification to its MCP connection. - async fn handle_mcp_over_acp_notification( - &mut self, - notification: Protocol::MessageNotification, - ) -> Result, crate::Error> { - let connection_id = Protocol::message_notification_connection_id(¬ification); - let Some(mcp_server_tx) = self.connections.get_mut(&connection_id) else { - return Ok(Handled::No { - message: notification, - retry: false, - }); - }; - let message = Protocol::into_message_notification(notification); - - let untyped = UntypedMessage { - method: message.method, - params: native_params_into_value(message.params), - }; - mcp_server_tx - .send(Dispatch::Notification(untyped)) - .await - .map_err(crate::Error::into_internal_error)?; - Ok(Handled::Yes) - } - - /// Disconnect an active native MCP-over-ACP connection. - fn handle_mcp_disconnect_request( - &mut self, - request: Protocol::DisconnectRequest, - responder: Responder, - ) -> Result< - Handled<( - Protocol::DisconnectRequest, - Responder, - )>, - crate::Error, - > { - let connection_id = Protocol::disconnect_connection_id(&request); - if self.connections.remove(&connection_id).is_none() { - return Ok(Handled::No { - message: (request, responder), - retry: false, + if let Some(service) = self.service.clone() { + let metadata = params + .as_ref() + .and_then(|params| params.get("_meta")) + .and_then(Value::as_object) + .expect("validated MCP metadata") + .clone(); + let cancellation = responder.cancellation(); + let operation_cancellation = McpOperationCancellation::new(); + let alive = Arc::new(futures::lock::Mutex::new(true)); + let send_connection = connection.clone(); + let send_server_id = server_id.clone(); + let send_request_id = request_id.clone(); + let send_cancellation = cancellation.clone(); + let send_operation_cancellation = operation_cancellation.clone(); + let send_alive = alive.clone(); + let notify = Arc::new(move |method: String, params: Option>| { + let connection = send_connection.clone(); + let server_id = send_server_id.clone(); + let request_id = send_request_id.clone(); + let cancellation = send_cancellation.clone(); + let operation_cancellation = send_operation_cancellation.clone(); + let alive = send_alive.clone(); + let send = async move { + let active = alive.lock().await; + if !*active + || cancellation.is_cancelled() + || operation_cancellation.is_cancelled() + { + return Err(crate::Error::request_cancelled()); + } + check_payload_size(&(&method, ¶ms), MAX_PAYLOAD_BYTES)?; + let send = connection.send_notification_to_async( + Agent, + Protocol::notification(server_id, request_id, method, params), + ); + futures::pin_mut!(send); + let cancelled = async { + let peer = cancellation.cancelled(); + let operation = operation_cancellation.cancelled(); + let shutdown = connection.shutdown_requested(); + futures::pin_mut!(peer, operation, shutdown); + let peer_or_operation = future::select(peer, operation); + futures::pin_mut!(peer_or_operation); + let _reason = future::select(peer_or_operation, shutdown).await; + }; + futures::pin_mut!(cancelled); + let result = match future::select(send, cancelled).await { + Either::Left((result, _)) => result, + Either::Right(((), _)) => Err(crate::Error::request_cancelled()), + }; + drop(active); + result + }; + Box::pin(send) as futures::future::BoxFuture<'static, Result<(), crate::Error>> }); + let context = McpRequestContext::new( + server_id.clone(), + request_id.clone(), + McpConnectionTo { + context: McpConnectionContext::Acp { + server_id, + request_id, + }, + connection: connection.clone(), + cleanup: Some(Arc::default()), + }, + metadata, + cancellation.clone(), + operation_cancellation.clone(), + notify, + ); + let is_discovery = method == "server/discover"; + let shutdown_connection = connection.clone(); + connection.spawn_protected(async move { + let request = McpRequest { method, params }; + let cleanup_connection = context.connection().clone(); + let operation = service.execute(request, context); + let stop = async { + let cancelled = cancellation.cancelled(); + let shutdown = shutdown_connection.shutdown_requested(); + futures::pin_mut!(cancelled); + futures::pin_mut!(shutdown); + let stop_rx = stop_rx; + futures::pin_mut!(stop_rx); + let cancel_or_shutdown = future::select(cancelled, shutdown); + futures::pin_mut!(cancel_or_shutdown); + let _reason = future::select(cancel_or_shutdown, stop_rx).await; + }; + let result = match future::select(operation, Box::pin(stop)).await { + Either::Left((result, _)) => result, + Either::Right(((), operation)) => { + operation_cancellation.cancel(); + *alive.lock().await = false; + // Do not discard the operation future: its completion + // includes rmcp handler cancellation and actor join. + drop(operation.await); + Err(crate::Error::request_cancelled()) + } + }; + *alive.lock().await = false; + cleanup_connection.wait_cleanup().await; + // Operation futures have been dropped and cannot send late output. + drop(guard); + let response = send_outcome::(responder, result, is_discovery); + if let Err(error) = response { + tracing::debug!(?error, "cannot send MCP response"); + } + Ok(()) + })?; + return Ok(Handled::Yes); } - responder.respond(Protocol::disconnect_response())?; + let cleanup_connection = McpConnectionTo { + context: McpConnectionContext::Acp { + server_id: server_id.clone(), + request_id: request_id.clone(), + }, + connection: connection.clone(), + cleanup: Some(Arc::default()), + }; + let connector = self.mcp_connect.clone(); + let connection_for_task = connection.clone(); + let cancellation = responder.cancellation(); + // Admission is atomic: this one protected task owns construction, + // execution, forwarding, and cleanup. A rejected spawn cannot start + // a backend or leave half of the operation running. + connection.spawn_protected(async move { + let backend = connector.connect(cleanup_connection.clone()); + let (mut client, server) = Channel::duplex(); + let mut backend = Some(Box::pin(backend.connect_to(server))); + let mut backend_error = None; + let inner_id = RequestId::Str(request_id.0.to_string()); + let is_discovery = method == "server/discover"; + let process = async { + let raw = RawJsonRpcMessage::request( + method, + params.map_or(Value::Null, Value::Object), + inner_id.clone(), + )?; + client + .tx + .send_frame(TransportFrame::Single(raw)) + .await + .map_err(crate::Error::into_internal_error)?; + loop { + let message = match backend.as_mut() { + Some(run) => match future::select(client.rx.next(), run).await { + Either::Left((message, _)) => message, + Either::Right((result, receive)) => { + drop(receive); + backend.take(); + backend_error = result.err(); + // Drain accepted output after backend exit, + // without waiting for an escaped sender to + // close or accepting any later output. + client.rx.close(); + continue; + } + }, + None => client.rx.next().await, + }; + let Some(budgeted) = message else { + return Err(backend_error.take().unwrap_or_else(|| { + crate::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend closed without a response", + ) + })); + }; + let (frame, _permit) = budgeted.into_parts(); + let TransportFrame::Single(message) = frame else { + return Err(crate::Error::new( + MCP_BACKEND_FAILURE, + "MCP backends must send individual valid JSON-RPC messages", + )); + }; + if matches!(message, RawJsonRpcMessage::Response(_)) + && message.response_id() != Some(&inner_id) + { + return Err(crate::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend returned a different request ID", + )); + } + match message { + RawJsonRpcMessage::Response(response) => { + // Returning ends notification forwarding before the terminal reply. + return match response { + RawJsonRpcResponse::Result { result, .. } => { + Ok(McpOutcome::Result(result)) + } + RawJsonRpcResponse::Error { error, .. } => { + Ok(McpOutcome::Error(into_mcp_error(*error))) + } + }; + } + RawJsonRpcMessage::Notification(notification) => { + check_payload_size(¬ification, MAX_PAYLOAD_BYTES)?; + let params = match notification.params { + Some(params) => match params.into_value() { + Value::Object(map) => Some(map), + _ => { + return Err(crate::Error::new( + MCP_BACKEND_FAILURE, + "MCP backend notification parameters must be an object", + )); + } + }, + None => None, + }; + connection_for_task + .send_notification_to_async( + Agent, + Protocol::notification( + server_id.clone(), + request_id.clone(), + notification.method.to_string(), + params, + ), + ) + .await?; + } + RawJsonRpcMessage::Request(_) => { + return Err(crate::Error::new( + MCP_BACKEND_FAILURE, + "reverse MCP requests are not supported", + )); + } + } + } + }; + let result = cancellation + .run_until_cancelled(async { + let process = process; + futures::pin_mut!(process); + let stop = async { + let _reason = future::select( + stop_rx, + Box::pin(connection_for_task.shutdown_requested()), + ) + .await; + }; + futures::pin_mut!(stop); + match future::select(process, stop).await { + Either::Left((result, _)) => result, + Either::Right(((), _)) => Err(crate::Error::request_cancelled()), + } + }) + .await; + // Drop the owned driver and revoke its channels before joining + // scoped tool cleanup, releasing the logical ID, or replying. + drop(backend); + drop(client); + cleanup_connection.wait_cleanup().await; + drop(guard); + let response = send_outcome::(responder, result, is_discovery); + if let Err(error) = response { + tracing::debug!(?error, "cannot send request-scoped MCP response"); + } + Ok(()) + })?; Ok(Handled::Yes) } } -impl HandleDispatchFrom +impl HandleDispatchFrom for McpActiveSession where Counterpart: HasPeer, - Protocol: McpProtocol, { fn describe_chain(&self) -> impl std::fmt::Debug { - "McpServerSession" + "McpServerRequests" } async fn handle_dispatch_from( @@ -450,31 +585,10 @@ where connection: ConnectionTo, ) -> Result, crate::Error> { MatchDispatchFrom::new(message, &connection) - .if_request_from( - Agent, - async |request: Protocol::ConnectRequest, responder| { - self.handle_connect_request(request, responder, &connection) - }, - ) - .await .if_request_from( Agent, async |request: Protocol::MessageRequest, responder| { - self.handle_mcp_over_acp_request(request, responder).await - }, - ) - .await - .if_notification_from( - Agent, - async |notification: Protocol::MessageNotification| { - self.handle_mcp_over_acp_notification(notification).await - }, - ) - .await - .if_request_from( - Agent, - async |request: Protocol::DisconnectRequest, responder| { - self.handle_mcp_disconnect_request(request, responder) + self.handle_request(request, responder, &connection) }, ) .await @@ -482,44 +596,237 @@ where } } -fn into_native_params(params: Value) -> Result>, crate::Error> { - match params { - Value::Null => Ok(None), - Value::Object(params) => Ok(Some(params)), - Value::Array(_) => Err(crate::Error::invalid_params() - .data("MCP-over-ACP only supports named inner MCP parameters")), - _ => { - Err(crate::Error::invalid_params() - .data("inner MCP parameters must be an object or null")) - } +/// Discovery describes the revisions available through this binding, not other +/// transports the hosted backend might also implement. +fn constrain_discovery_versions(result: &mut Value) -> Result<(), crate::Error> { + let versions = result + .get_mut("supportedVersions") + .and_then(Value::as_array_mut) + .ok_or_else(|| crate::Error::internal_error().data("invalid MCP discovery result"))?; + if !versions + .iter() + .any(|version| version.as_str() == Some(MCP_VERSION)) + { + return Err(crate::Error::new(-32022, "Unsupported protocol version") + .data(serde_json::json!({"requested": MCP_VERSION, "supported": versions}))); } + *versions = vec![Value::String(MCP_VERSION.to_owned())]; + Ok(()) } -fn native_params_into_value(params: Option>) -> Value { - params.map_or(Value::Null, Value::Object) +fn validate_modern_request( + method: &str, + params: Option<&Map>, +) -> Result<(), crate::Error> { + if method == "initialize" { + return Err( + crate::Error::method_not_found().data("native MCP requests do not use initialize") + ); + } + let meta = params + .and_then(|params| params.get("_meta")) + .and_then(Value::as_object) + .ok_or_else(|| { + crate::Error::invalid_params().data("inner params._meta must be an object") + })?; + let version = meta + .get("io.modelcontextprotocol/protocolVersion") + .and_then(Value::as_str) + .ok_or_else(|| { + crate::Error::invalid_params() + .data("inner params._meta requires io.modelcontextprotocol/protocolVersion") + })?; + if version != MCP_VERSION { + return Err(crate::Error::new(-32022, "Unsupported protocol version") + .data(serde_json::json!({"requested": version, "supported": [MCP_VERSION]}))); + } + if !meta + .get("io.modelcontextprotocol/clientCapabilities") + .is_some_and(Value::is_object) + { + return Err(crate::Error::invalid_params().data( + "inner params._meta requires io.modelcontextprotocol/clientCapabilities object", + )); + } + Ok(()) } #[cfg(test)] mod tests { + use super::{ + ActiveRequests, MAX_ACTIVE_REQUESTS, MAX_PAYLOAD_BYTES, McpOutcome, V1McpProtocol, + admit_request, check_payload_size, constrain_discovery_versions, into_mcp_error, + outcome_response, validate_modern_request, + }; + use crate::{ + mcp_server::MCP_RESOURCE_EXHAUSTED, + schema::v1::{McpError, McpRequestId}, + }; use serde_json::json; - use super::{into_native_params, native_params_into_value}; + #[test] + fn both_outcome_branches_obey_the_binding_payload_limit() { + for outcome in [ + McpOutcome::Result(json!("x".repeat(MAX_PAYLOAD_BYTES))), + McpOutcome::Error( + McpError::new(-32000, "peer error").data(json!("x".repeat(MAX_PAYLOAD_BYTES))), + ), + ] { + let error = outcome_response::(outcome) + .expect_err("oversized carrier must be rejected"); + assert_eq!(i32::from(error.code), MCP_RESOURCE_EXHAUSTED); + } + let result = + outcome_response::(McpOutcome::Result(serde_json::Value::Null)).unwrap(); + assert_eq!( + serde_json::to_value(result).unwrap(), + json!({"result":null}) + ); + } + #[cfg(feature = "unstable_protocol_v2")] #[test] - fn native_mcp_params_round_trip_objects_and_null() { - let object = json!({ "name": "echo", "arguments": {} }); - let params = into_native_params(object.clone()).expect("object params should be valid"); - assert_eq!(native_params_into_value(params), object); + fn versioned_outcomes_preserve_results_and_error_fields() { + for value in [ + serde_json::Value::Null, + json!({"resultType":"complete","_meta":{"custom":true}}), + ] { + let v1 = outcome_response::(McpOutcome::Result(value.clone())).unwrap(); + let v2 = outcome_response::(McpOutcome::Result(value)).unwrap(); + assert_eq!( + serde_json::to_value(v1).unwrap(), + serde_json::to_value(v2).unwrap() + ); + } + for data in [ + None, + Some(serde_json::Value::Null), + Some(json!({"details":[1,2]})), + ] { + let mut error = McpError::new(-32000, "opaque peer error"); + if let Some(data) = data { + error = error.data(data); + } + error + .extra + .insert("extension".into(), json!({"preserve":true})); + let expected = json!({"error": error}); + let v2: crate::schema::v2::MessageMcpResponse = + outcome_response::(McpOutcome::Error(error)).unwrap(); + assert_eq!(serde_json::to_value(v2).unwrap(), expected); + } + } - let params = into_native_params(serde_json::Value::Null) - .expect("omitted params should be represented as null"); - assert_eq!(native_params_into_value(params), serde_json::Value::Null); + #[test] + fn only_modern_request_metadata_is_accepted() { + let modern = json!({"_meta": {"io.modelcontextprotocol/protocolVersion": "2026-07-28", "io.modelcontextprotocol/clientCapabilities": {}, "requestState": {"opaque": true}}}); + assert!(validate_modern_request("tools/list", modern.as_object()).is_ok()); + assert!(validate_modern_request("initialize", modern.as_object()).is_err()); + assert!(validate_modern_request("tools/list", json!({"_meta": {"io.modelcontextprotocol/protocolVersion": "2025-03-26", "io.modelcontextprotocol/clientCapabilities": {}}}).as_object()).is_err()); + assert!( + validate_modern_request( + "tools/list", + json!({"_meta": {"io.modelcontextprotocol/protocolVersion": "2026-07-28"}}) + .as_object() + ) + .is_err() + ); + } + + #[test] + fn unsupported_version_is_an_mcp_error_not_a_legacy_fallback() { + let params = json!({ + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2025-11-25", + "io.modelcontextprotocol/clientCapabilities": {} + } + }); + let error = validate_modern_request("tools/call", params.as_object()).unwrap_err(); + assert_eq!( + serde_json::to_value(error).unwrap(), + json!({ + "code": -32022, + "message": "Unsupported protocol version", + "data": {"requested": "2025-11-25", "supported": ["2026-07-28"]} + }) + ); + } + + #[test] + fn native_request_admission_is_bounded_and_recovers_after_cleanup() { + let active = ActiveRequests::default(); + let mut admitted = Vec::new(); + for index in 0..MAX_ACTIVE_REQUESTS { + admitted + .push(admit_request(&active, McpRequestId::new(format!("req-{index}"))).unwrap()); + } + assert_eq!(active.lock().unwrap().len(), MAX_ACTIVE_REQUESTS); + let duplicate = admit_request(&active, McpRequestId::new("req-0")) + .err() + .unwrap(); + assert_eq!(duplicate.code, crate::ErrorCode::InvalidParams); + let overload = admit_request(&active, McpRequestId::new("extra")) + .err() + .unwrap(); + assert_eq!(i32::from(overload.code), MCP_RESOURCE_EXHAUSTED); + drop(admitted.pop()); + let replacement = admit_request(&active, McpRequestId::new("replacement")).unwrap(); + assert_eq!(active.lock().unwrap().len(), MAX_ACTIVE_REQUESTS); + drop(replacement); + drop(admitted); + assert!(active.lock().unwrap().is_empty()); + } + + #[test] + fn payload_limits_count_json_escaping_without_building_an_extra_buffer() { + let payload = json!({"text": "\n\n"}); + let encoded = serde_json::to_vec(&payload).unwrap(); + assert!(check_payload_size(&payload, encoded.len()).is_ok()); + assert!(check_payload_size(&payload, encoded.len() - 1).is_err()); + } + + #[test] + fn discovery_reports_the_binding_version_without_changing_other_payload() { + let mut result = json!({ + "resultType": "complete", + "supportedVersions": ["2025-11-25", "2026-07-28"], + "capabilities": {"tools": {}}, + "_meta": {"vendor/opaque": ["preserved"]} + }); + constrain_discovery_versions(&mut result).unwrap(); + assert_eq!(result["supportedVersions"], json!(["2026-07-28"])); + assert_eq!(result["_meta"]["vendor/opaque"], json!(["preserved"])); + assert_eq!(result["capabilities"], json!({"tools": {}})); + let mut unsupported = json!({"supportedVersions": ["2025-11-25"]}); + assert!(constrain_discovery_versions(&mut unsupported).is_err()); + assert!(constrain_discovery_versions(&mut json!({})).is_err()); } #[test] - fn native_mcp_params_reject_positional_params() { - let error = into_native_params(json!(["positional"])) - .expect_err("native MCP-over-ACP cannot represent positional params"); - assert_eq!(error.code, crate::ErrorCode::InvalidParams); + fn mcp_validation_errors_use_inner_carrier_and_preserve_null_data() { + let unsupported = validate_modern_request( + "tools/list", + json!({"_meta": { + "io.modelcontextprotocol/protocolVersion": "2025-03-26", + "io.modelcontextprotocol/clientCapabilities": {} + }}) + .as_object(), + ) + .expect_err("unsupported inner version"); + let response = + outcome_response::(McpOutcome::Error(into_mcp_error(unsupported))) + .unwrap(); + assert_eq!( + serde_json::to_value(response).unwrap()["error"]["code"], + -32022 + ); + let response = outcome_response::(McpOutcome::Error( + McpError::new(-32000, "opaque MCP error").data(serde_json::Value::Null), + )) + .unwrap(); + assert_eq!( + serde_json::to_value(response).unwrap(), + json!({"error": {"code": -32000, "message": "opaque MCP error", "data": null}}) + ); } } diff --git a/src/agent-client-protocol/src/mcp_server/context.rs b/src/agent-client-protocol/src/mcp_server/context.rs index aba4b938..cb2f3455 100644 --- a/src/agent-client-protocol/src/mcp_server/context.rs +++ b/src/agent-client-protocol/src/mcp_server/context.rs @@ -1,7 +1,11 @@ use crate::{ConnectionTo, role::Role}; +#[cfg(feature = "unstable_mcp_over_acp")] +use futures::channel::oneshot; +#[cfg(feature = "unstable_mcp_over_acp")] +use std::sync::{Arc, Mutex}; #[cfg(feature = "unstable_mcp_over_acp")] -use crate::schema::v1::{McpConnectionId, McpServerAcpId}; +use crate::schema::v1::{McpRequestId, McpServerAcpId}; /// Describes how an MCP server connection was established. #[derive(Clone, Debug, PartialEq, Eq)] @@ -16,8 +20,8 @@ pub enum McpConnectionContext { /// The identifier advertised in the session's `McpServer::Acp` declaration. server_id: McpServerAcpId, - /// The identifier for this active `mcp/connect` connection. - connection_id: McpConnectionId, + /// The logical identifier of this independent MCP request. + request_id: McpRequestId, }, } @@ -40,15 +44,15 @@ impl McpConnectionContext { } } - /// The identifier for the active `mcp/connect` connection. + /// The logical identifier of the active MCP request. /// /// Returns `None` for a standalone MCP connection. #[cfg(feature = "unstable_mcp_over_acp")] #[must_use] - pub fn connection_id(&self) -> Option<&McpConnectionId> { + pub fn request_id(&self) -> Option<&McpRequestId> { match self { Self::Standalone => None, - Self::Acp { connection_id, .. } => Some(connection_id), + Self::Acp { request_id, .. } => Some(request_id), } } } @@ -58,9 +62,28 @@ impl McpConnectionContext { pub struct McpConnectionTo { pub(super) context: McpConnectionContext, pub(super) connection: ConnectionTo, + #[cfg(feature = "unstable_mcp_over_acp")] + pub(super) cleanup: Option>>>>, } impl McpConnectionTo { + #[cfg(all(feature = "unstable_mcp_over_acp", feature = "schemars"))] + pub(crate) fn register_cleanup(&self, done: oneshot::Receiver<()>) { + if let Some(cleanup) = &self.cleanup { + cleanup.lock().expect("MCP cleanup poisoned").push(done); + } + } + + #[cfg(feature = "unstable_mcp_over_acp")] + pub(crate) async fn wait_cleanup(&self) { + if let Some(cleanup) = &self.cleanup { + let pending = std::mem::take(&mut *cleanup.lock().expect("MCP cleanup poisoned")); + for done in pending { + let _ = done.await; + } + } + } + /// Describes whether this is a standalone or ACP-attached MCP connection. #[must_use] pub fn context(&self) -> &McpConnectionContext { @@ -76,13 +99,13 @@ impl McpConnectionTo { self.context.server_id() } - /// The identifier for the active `mcp/connect` connection. + /// The logical identifier of the active MCP request. /// /// Returns `None` for a standalone MCP connection. #[cfg(feature = "unstable_mcp_over_acp")] #[must_use] - pub fn connection_id(&self) -> Option<&McpConnectionId> { - self.context.connection_id() + pub fn request_id(&self) -> Option<&McpRequestId> { + self.context.request_id() } /// Borrow the host protocol connection. @@ -108,24 +131,24 @@ mod tests { #[cfg(feature = "unstable_mcp_over_acp")] { assert_eq!(context.server_id(), None); - assert_eq!(context.connection_id(), None); + assert_eq!(context.request_id(), None); } } #[cfg(feature = "unstable_mcp_over_acp")] #[test] - fn acp_context_exposes_server_and_connection_ids() { - use crate::schema::v1::{McpConnectionId, McpServerAcpId}; + fn acp_context_exposes_server_and_request_ids() { + use crate::schema::v1::{McpRequestId, McpServerAcpId}; let server_id = McpServerAcpId::new("server-id"); - let connection_id = McpConnectionId::new("connection-id"); + let request_id = McpRequestId::new("request-id"); let context = McpConnectionContext::Acp { server_id: server_id.clone(), - connection_id: connection_id.clone(), + request_id: request_id.clone(), }; assert!(!context.is_standalone()); assert_eq!(context.server_id(), Some(&server_id)); - assert_eq!(context.connection_id(), Some(&connection_id)); + assert_eq!(context.request_id(), Some(&request_id)); } } diff --git a/src/agent-client-protocol/src/mcp_server/mod.rs b/src/agent-client-protocol/src/mcp_server/mod.rs index a77b53a1..e19d2d94 100644 --- a/src/agent-client-protocol/src/mcp_server/mod.rs +++ b/src/agent-client-protocol/src/mcp_server/mod.rs @@ -57,6 +57,8 @@ mod context; #[cfg(feature = "schemars")] mod registry; mod server; +#[cfg(feature = "unstable_mcp_over_acp")] +mod service; #[cfg(feature = "schemars")] mod tool; #[cfg(feature = "schemars")] @@ -70,9 +72,20 @@ pub use registry::{ EnabledTools, McpToolMetadata, McpToolRegistry, McpToolSchema, RegisteredMcpTool, }; pub use server::McpServer; +#[cfg(feature = "unstable_mcp_over_acp")] +pub use service::{ + McpOperationCancellation, McpOutcome, McpRequest, McpRequestContext, McpService, +}; #[cfg(feature = "schemars")] #[cfg_attr(docsrs, doc(cfg(feature = "schemars")))] pub use tool::McpTool; #[cfg(feature = "schemars")] #[cfg_attr(docsrs, doc(cfg(feature = "schemars")))] pub use tool_fn::{tool_fn, tool_fn_mut}; + +/// ACP binding error: the MCP operation admission or payload budget was exhausted. +pub const MCP_RESOURCE_EXHAUSTED: i32 = -33000; +/// ACP binding error: the requested MCP server registration is no longer available. +pub const MCP_SERVER_UNAVAILABLE: i32 = -33001; +/// ACP binding error: the MCP backend failed before returning an MCP outcome. +pub const MCP_BACKEND_FAILURE: i32 = -33002; diff --git a/src/agent-client-protocol/src/mcp_server/server.rs b/src/agent-client-protocol/src/mcp_server/server.rs index 2b34859f..5f86ee18 100644 --- a/src/agent-client-protocol/src/mcp_server/server.rs +++ b/src/agent-client-protocol/src/mcp_server/server.rs @@ -4,6 +4,8 @@ use std::{marker::PhantomData, sync::Arc}; use futures::{StreamExt, channel::mpsc}; +#[cfg(feature = "unstable_mcp_over_acp")] +use crate::mcp_server::McpService; use crate::{ ConnectTo, Dispatch, DynConnectTo, Role, jsonrpc::run::{NullRun, RunWithConnectionTo}, @@ -66,6 +68,8 @@ pub struct McpServer { /// The "connect" instance connect: Arc>, + #[cfg(feature = "unstable_mcp_over_acp")] + service: Option>>, /// The runner is a task that should be run alongside the message handler. /// Some futures direct messages back through channels to this future which actually @@ -103,6 +107,37 @@ where McpServer { phantom: PhantomData, connect: Arc::new(c), + #[cfg(feature = "unstable_mcp_over_acp")] + service: None, + runner, + } + } + + /// Construct a reusable request-native application service for ACP. + /// + /// Standalone serving additionally needs a direct MCP transport adapter; + /// use [`Self::new_service_with_standalone`] when direct serving is needed. + #[cfg(feature = "unstable_mcp_over_acp")] + pub fn new_service( + service: impl McpService, + name: impl Into, + runner: Run, + ) -> Self { + Self::new_service_with_standalone(service, NoStandalone { name: name.into() }, runner) + } + + /// Construct a reusable application service with a separate standalone MCP + /// connector. ACP requests never create connector sessions. + #[cfg(feature = "unstable_mcp_over_acp")] + pub fn new_service_with_standalone( + service: impl McpService, + standalone: impl McpServerConnect, + runner: Run, + ) -> Self { + Self { + phantom: PhantomData, + connect: Arc::new(standalone), + service: Some(Arc::new(service)), runner, } } @@ -116,10 +151,14 @@ where let Self { phantom: _, connect, + service, runner, } = self; let server_id = McpServerAcpId::new(format!("mcp-server:{}", Uuid::new_v4())); - (McpSessionHandler::new(server_id, connect), runner) + ( + McpSessionHandler::new_with_service(server_id, connect, service), + runner, + ) } /// Split this MCP server into a protocol v2 session handler and its runner. @@ -131,10 +170,40 @@ where let Self { phantom: _, connect, + service, runner, } = self; let server_id = McpServerAcpId::new(format!("mcp-server:{}", Uuid::new_v4())); - (V2McpSessionHandler::new(server_id, connect), runner) + ( + V2McpSessionHandler::new_with_service(server_id, connect, service), + runner, + ) + } +} + +#[cfg(feature = "unstable_mcp_over_acp")] +struct NoStandalone { + name: String, +} + +#[cfg(feature = "unstable_mcp_over_acp")] +impl McpServerConnect for NoStandalone { + fn name(&self) -> String { + self.name.clone() + } + + fn connect(&self, _context: McpConnectionTo) -> DynConnectTo { + struct Unavailable; + impl ConnectTo for Unavailable { + fn connect_to( + self, + _client: impl ConnectTo, + ) -> impl Future> + Send { + std::future::ready(Err(crate::Error::method_not_found() + .data("this MCP service has no standalone transport adapter"))) + } + } + DynConnectTo::new(Unavailable) } } @@ -154,9 +223,17 @@ impl McpSessionHandler where Counterpart: HasPeer, { - pub fn new(server_id: McpServerAcpId, connect: Arc>) -> Self { + fn new_with_service( + server_id: McpServerAcpId, + connect: Arc>, + service: Option>>, + ) -> Self { Self { - active_session: McpActiveSession::new(server_id.clone(), connect.clone()), + active_session: McpActiveSession::new_with_service( + server_id.clone(), + connect.clone(), + service, + ), server_id, connect, } @@ -187,9 +264,22 @@ impl V2McpSessionHandler where Counterpart: HasPeer, { + #[cfg(test)] fn new(server_id: McpServerAcpId, connect: Arc>) -> Self { + Self::new_with_service(server_id, connect, None) + } + + fn new_with_service( + server_id: McpServerAcpId, + connect: Arc>, + service: Option>>, + ) -> Self { Self { - active_session: McpActiveSession::new(server_id.clone(), connect.clone()), + active_session: McpActiveSession::new_with_service( + server_id.clone(), + connect.clone(), + service, + ), server_id, connect, } @@ -405,6 +495,8 @@ where connect, runner, phantom: _, + #[cfg(feature = "unstable_mcp_over_acp")] + service: _, } = self; let (tx, mut rx) = mpsc::unbounded(); @@ -424,6 +516,8 @@ where connect.connect(McpConnectionTo { context: McpConnectionContext::Standalone, connection: connection_to_client.clone(), + #[cfg(feature = "unstable_mcp_over_acp")] + cleanup: None, }); role::mcp::Client diff --git a/src/agent-client-protocol/src/mcp_server/service.rs b/src/agent-client-protocol/src/mcp_server/service.rs new file mode 100644 index 00000000..30eee4a5 --- /dev/null +++ b/src/agent-client-protocol/src/mcp_server/service.rs @@ -0,0 +1,197 @@ +//! Request-native application services for MCP-over-ACP. + +use std::sync::Arc; + +use futures::{ + channel::oneshot, + future::{BoxFuture, FutureExt, Shared}, +}; +use serde_json::{Map, Value}; + +use super::McpConnectionTo; +use crate::{ + Error, RequestCancellation, Role, + schema::v1::{McpError, McpRequestId, McpServerAcpId}, +}; + +/// One MCP invocation. Its application service may be reused across invocations. +#[derive(Debug)] +pub struct McpRequest { + /// The MCP method. + pub method: String, + /// Its MCP parameters; metadata is validated before dispatch. + pub params: Option>, +} + +/// The MCP outcome is distinct from a failure in the ACP binding itself. +#[derive(Debug)] +pub enum McpOutcome { + /// Successful, opaque MCP result. + Result(Value), + /// An unmodified MCP error object (including optional or explicitly null data). + Error(McpError), +} + +type Notify = dyn Fn(String, Option>) -> BoxFuture<'static, Result<(), Error>> + + Send + + Sync; + +/// Explicit cancellation of an operation, including provider removal and +/// connection shutdown (which need not cancel the original ACP request). +#[derive(Clone)] +pub struct McpOperationCancellation { + state: Arc, +} + +impl std::fmt::Debug for McpOperationCancellation { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("McpOperationCancellation") + .field("cancelled", &self.is_cancelled()) + .finish() + } +} + +struct CancellationState { + cancelled: std::sync::atomic::AtomicBool, + sender: std::sync::Mutex>>, + signal: Shared>, +} + +impl McpOperationCancellation { + pub(crate) fn new() -> Self { + let (tx, rx) = oneshot::channel(); + Self { + state: Arc::new(CancellationState { + cancelled: std::sync::atomic::AtomicBool::new(false), + sender: std::sync::Mutex::new(Some(tx)), + signal: rx.map(|_| ()).boxed().shared(), + }), + } + } + + pub(crate) fn cancel(&self) { + self.state + .cancelled + .store(true, std::sync::atomic::Ordering::Release); + drop( + self.state + .sender + .lock() + .expect("MCP cancellation poisoned") + .take(), + ); + } + + /// Await cancellation from the caller, provider, or transport. + pub async fn cancelled(&self) { + self.state.signal.clone().await; + } + /// Whether the operation may still produce output. + #[must_use] + pub fn is_cancelled(&self) -> bool { + self.state + .cancelled + .load(std::sync::atomic::Ordering::Acquire) + } +} + +/// Per-operation authority. Notifications are admitted only while this request +/// is live; retaining the service does not retain an operation's output rights. +#[derive(Clone)] +pub struct McpRequestContext { + server_id: McpServerAcpId, + request_id: McpRequestId, + connection: McpConnectionTo, + metadata: Map, + cancellation: RequestCancellation, + operation_cancellation: McpOperationCancellation, + notify: Arc, +} + +impl std::fmt::Debug for McpRequestContext { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("McpRequestContext") + .field("server_id", &self.server_id) + .field("request_id", &self.request_id) + .field("metadata", &self.metadata) + .field("operation_cancellation", &self.operation_cancellation) + .finish_non_exhaustive() + } +} + +impl McpRequestContext { + pub(crate) fn new( + server_id: McpServerAcpId, + request_id: McpRequestId, + connection: McpConnectionTo, + metadata: Map, + cancellation: RequestCancellation, + operation_cancellation: McpOperationCancellation, + notify: Arc, + ) -> Self { + Self { + server_id, + request_id, + connection, + metadata, + cancellation, + operation_cancellation, + notify, + } + } + + /// Server identifier bound to this operation. + pub fn server_id(&self) -> &McpServerAcpId { + &self.server_id + } + /// Logical operation identifier. + pub fn request_id(&self) -> &McpRequestId { + &self.request_id + } + /// Host connection, available to application tools. + pub fn connection(&self) -> &McpConnectionTo { + &self.connection + } + /// Validated MCP metadata, including the negotiated protocol version and + /// the client's capability declaration. + pub fn metadata(&self) -> &Map { + &self.metadata + } + /// Request cancellation handle. + pub fn cancellation(&self) -> &RequestCancellation { + &self.cancellation + } + /// Cancellation for this operation, including provider removal and EOF. + pub fn operation_cancellation(&self) -> &McpOperationCancellation { + &self.operation_cancellation + } + + /// Send a bounded, operation-scoped MCP notification. + pub async fn send_notification( + &self, + method: impl Into, + params: Option>, + ) -> Result<(), Error> { + if self.cancellation.is_cancelled() || self.operation_cancellation.is_cancelled() { + return Err(Error::request_cancelled()); + } + (self.notify)(method.into(), params).await + } +} + +/// Reusable application service. An invocation owns its returned future; an +/// implementation may deliberately share application state between requests. +pub trait McpService: Send + Sync + 'static { + /// Execute one MCP request, returning an owned operation future. + /// + /// The future includes backend teardown: on + /// [`McpRequestContext::operation_cancellation`], stop user work and finish + /// owned cleanup before returning. The binding keeps admission until this + /// future completes rather than abandoning cleanup by dropping it. The rmcp + /// adapter implements this supervision for its handler futures. + fn execute( + &self, + request: McpRequest, + context: McpRequestContext, + ) -> BoxFuture<'static, Result>; +} diff --git a/src/agent-client-protocol/src/mcp_server/tool_fn.rs b/src/agent-client-protocol/src/mcp_server/tool_fn.rs index 2fa645d4..79a8e569 100644 --- a/src/agent-client-protocol/src/mcp_server/tool_fn.rs +++ b/src/agent-client-protocol/src/mcp_server/tool_fn.rs @@ -1,12 +1,13 @@ //! Runtime-neutral helpers for registering function-backed MCP tools. use futures::{ - SinkExt, StreamExt, - channel::{mpsc, oneshot}, - future::BoxFuture, + StreamExt, + channel::oneshot, + future::{self, BoxFuture, Either}, }; use schemars::JsonSchema; use serde::{Serialize, de::DeserializeOwned}; +use std::pin::Pin; use crate::{ConnectionTo, Error, Role, RunWithConnectionTo}; @@ -16,11 +17,12 @@ struct ToolCall { params: P, mcp_connection: McpConnectionTo, result_tx: futures::channel::oneshot::Sender>, + done_tx: oneshot::Sender<()>, } struct ToolFnMutRunner { func: F, - call_rx: mpsc::Receiver>, + call_rx: Pin>>>, tool_future_fn: Box< dyn for<'a> Fn( &'a mut F, @@ -52,13 +54,34 @@ where while let Some(ToolCall { params, mcp_connection, - result_tx, + mut result_tx, + done_tx, }) = call_rx.next().await { - let result = tool_future_fn(&mut func, params, mcp_connection).await; - result_tx - .send(result) - .map_err(|_| crate::util::internal_error("failed to send MCP result"))?; + // The caller may have cancelled while this invocation waited behind + // another mutable tool call. Do not start work for a gone caller. + if result_tx.is_canceled() { + drop(params); + drop(mcp_connection); + drop(result_tx); + let _ = done_tx.send(()); + continue; + } + let result = { + let cancelled = result_tx.cancellation(); + futures::pin_mut!(cancelled); + match future::select(tool_future_fn(&mut func, params, mcp_connection), cancelled) + .await + { + Either::Left((result, _)) => Some(result), + Either::Right(((), _)) => None, + } + }; + if let Some(result) = result { + // Cancellation after execution is not a runner failure. + drop(result_tx.send(result)); + } + let _ = done_tx.send(()); } Ok(()) } @@ -66,7 +89,7 @@ where struct ToolFnRunner { func: F, - call_rx: mpsc::Receiver>, + call_rx: Pin>>>, tool_future_fn: Box< dyn for<'a> Fn(&'a F, P, McpConnectionTo) -> BoxFuture<'a, Result> + Send @@ -92,9 +115,8 @@ where call_rx, tool_future_fn, } = self; - crate::util::process_stream_concurrently( - call_rx, - async |tool_call| { + call_rx + .for_each_concurrent(64, |tool_call| { fn hack<'a, F, P, R, MyRole>( func: &'a F, params: P, @@ -108,7 +130,8 @@ where + Send + Sync ), - result_tx: oneshot::Sender>, + mut result_tx: oneshot::Sender>, + done_tx: oneshot::Sender<()>, ) -> BoxFuture<'a, ()> where MyRole: Role, @@ -117,8 +140,30 @@ where F: Send + Sync, { Box::pin(async move { - let result = tool_future_fn(func, params, mcp_connection).await; - drop(result_tx.send(result)); + if result_tx.is_canceled() { + drop(params); + drop(mcp_connection); + drop(result_tx); + let _ = done_tx.send(()); + return; + } + let result = { + let cancelled = result_tx.cancellation(); + futures::pin_mut!(cancelled); + match future::select( + tool_future_fn(func, params, mcp_connection), + cancelled, + ) + .await + { + Either::Left((result, _)) => Some(result), + Either::Right(((), _)) => None, + } + }; + if let Some(result) = result { + drop(result_tx.send(result)); + } + let _ = done_tx.send(()); }) } @@ -126,21 +171,27 @@ where params, mcp_connection, result_tx, + done_tx, } = tool_call; - hack(&func, params, mcp_connection, &*tool_future_fn, result_tx).await; - Ok(()) - }, - |a, b| Box::pin(a(b)), - ) - .await + hack( + &func, + params, + mcp_connection, + &*tool_future_fn, + result_tx, + done_tx, + ) + }) + .await; + Ok(()) } } struct ToolFnTool { name: String, description: String, - call_tx: mpsc::Sender>, + call_tx: async_channel::Sender>, } impl McpTool for ToolFnTool @@ -162,13 +213,18 @@ where async fn call_tool(&self, params: P, mcp_connection: McpConnectionTo) -> Result { let (result_tx, result_rx) = oneshot::channel(); + let (done_tx, done_rx) = oneshot::channel(); + #[cfg(feature = "unstable_mcp_over_acp")] + mcp_connection.register_cleanup(done_rx); + #[cfg(not(feature = "unstable_mcp_over_acp"))] + let _done_rx = done_rx; self.call_tx - .clone() .send(ToolCall { params, mcp_connection, result_tx, + done_tx, }) .await .map_err(crate::util::internal_error)?; @@ -192,7 +248,7 @@ pub fn tool_fn_mut( + Send + 'static, ) -> ( - impl McpTool + 'static, + impl McpTool + 'static, impl RunWithConnectionTo, ) where @@ -201,7 +257,7 @@ where Ret: JsonSchema + Serialize + 'static + Send, F: AsyncFnMut(P, McpConnectionTo) -> Result + Send, { - let (call_tx, call_rx) = mpsc::channel(128); + let (call_tx, call_rx) = async_channel::bounded(128); ( ToolFnTool { name: name.to_string(), @@ -210,7 +266,7 @@ where }, ToolFnMutRunner { func, - call_rx, + call_rx: Box::pin(call_rx), tool_future_fn: Box::new(tool_future_fn), }, ) @@ -230,7 +286,7 @@ pub fn tool_fn( + Sync + 'static, ) -> ( - impl McpTool + 'static, + impl McpTool + 'static, impl RunWithConnectionTo, ) where @@ -239,7 +295,7 @@ where Ret: JsonSchema + Serialize + 'static + Send, F: AsyncFn(P, McpConnectionTo) -> Result + Send + Sync + 'static, { - let (call_tx, call_rx) = mpsc::channel(128); + let (call_tx, call_rx) = async_channel::bounded(128); ( ToolFnTool { name: name.to_string(), @@ -248,7 +304,7 @@ where }, ToolFnRunner { func, - call_rx, + call_rx: Box::pin(call_rx), tool_future_fn: Box::new(tool_future_fn), }, ) diff --git a/src/agent-client-protocol/src/role/acp.rs b/src/agent-client-protocol/src/role/acp.rs index f99c1ecc..067bc045 100644 --- a/src/agent-client-protocol/src/role/acp.rs +++ b/src/agent-client-protocol/src/role/acp.rs @@ -17,16 +17,16 @@ use crate::role::{HasPeer, RemoteStyle}; #[cfg(not(feature = "unstable_protocol_v2"))] use crate::schema::InitializeProxyRequest; use crate::schema::METHOD_INITIALIZE_PROXY; +#[cfg(feature = "unstable_protocol_v2")] +use crate::schema::v1::RequestId; use crate::schema::v1::{InitializeRequest, SessionId}; #[cfg(not(feature = "unstable_protocol_v2"))] use crate::schema::v1::{NewSessionRequest, NewSessionResponse}; #[cfg(feature = "unstable_protocol_v2")] -use crate::schema::v1::{RequestId, Response as RpcResponse}; -#[cfg(feature = "unstable_protocol_v2")] use crate::schema::{ProtocolVersion, v2}; use crate::util::MatchDispatchFrom; #[cfg(feature = "unstable_protocol_v2")] -use crate::{Channel, RawJsonRpcMessage, RawJsonRpcParams}; +use crate::{Channel, RawJsonRpcMessage, RawJsonRpcParams, RawJsonRpcResponse as RpcResponse}; use crate::{ConnectTo, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, Role, RoleId}; #[cfg(feature = "unstable_protocol_v2")] @@ -886,7 +886,7 @@ fn invalid_initialize_params(error: impl ToString) -> crate::Error { #[cfg(feature = "unstable_protocol_v2")] fn send_initialize_error( - tx: &futures::channel::mpsc::UnboundedSender, + tx: &crate::jsonrpc::FrameSender, frame: &TransportFrame, error: crate::Error, ) -> Result<(), crate::Error> { @@ -944,8 +944,7 @@ fn send_initialize_error( } }; - tx.unbounded_send(response) - .map_err(crate::util::internal_error) + tx.try_send(response).map_err(crate::util::internal_error) } #[cfg(feature = "unstable_protocol_v2")] @@ -972,8 +971,8 @@ async fn reject_initialize( #[cfg(feature = "unstable_protocol_v2")] struct RunningProtocolPeer { - rx: futures::channel::mpsc::UnboundedReceiver, - tx: futures::channel::mpsc::UnboundedSender, + rx: crate::jsonrpc::FrameReceiver, + tx: crate::jsonrpc::FrameSender, future: crate::BoxFuture<'static, Result<(), crate::Error>>, } @@ -981,14 +980,18 @@ struct RunningProtocolPeer { impl RunningProtocolPeer { fn new(component: impl ConnectTo) -> Self { let (Channel { rx, tx }, future) = component.into_channel_and_future(); - Self { rx, tx, future } + Self { + rx, + tx, + future: Box::pin(future), + } } async fn next_frame(self) -> Result, crate::Error> { let Self { mut rx, tx, future } = self; match future::select(Box::pin(rx.next()), future).await { future::Either::Left((Some(frame), future)) => { - Ok(Some((frame, Self { rx, tx, future }))) + Ok(Some((frame.into_frame(), Self { rx, tx, future }))) } future::Either::Left((None, future)) => { future.await?; @@ -1001,7 +1004,7 @@ impl RunningProtocolPeer { return Ok(None); }; Ok(Some(( - frame, + frame.into_frame(), Self { rx, tx, @@ -1024,9 +1027,7 @@ impl RunningProtocolPeer { } fn send_frame(&self, frame: TransportFrame) -> Result<(), crate::Error> { - self.tx - .unbounded_send(frame) - .map_err(crate::util::internal_error) + self.tx.try_send(frame).map_err(crate::util::internal_error) } } @@ -1133,7 +1134,7 @@ impl InitializeResponse { }), RawJsonRpcMessage::Response(RpcResponse::Error { id, error }) => Ok(Self { id, - result: Err(error), + result: Err(error.into_acp_error()), }), message => Err(crate::Error::invalid_request().data(format!( "first ACP response must be an initialize response, got {message:?}", diff --git a/src/agent-client-protocol/src/schema/enum_impls.rs b/src/agent-client-protocol/src/schema/enum_impls.rs index 8a937c43..e485e868 100644 --- a/src/agent-client-protocol/src/schema/enum_impls.rs +++ b/src/agent-client-protocol/src/schema/enum_impls.rs @@ -31,8 +31,6 @@ impl_jsonrpc_request_enum!(ClientRequest { SetSessionModeRequest => "session/set_mode", SetSessionConfigOptionRequest => "session/set_config_option", PromptRequest => "session/prompt", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpRequest => "mcp/message", [ext] ExtMethodRequest, }); @@ -57,8 +55,6 @@ impl_jsonrpc_response_enum!(AgentResponse { SetSessionModeResponse => "session/set_mode", SetSessionConfigOptionResponse => "session/set_config_option", PromptResponse => "session/prompt", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpResponse => "mcp/message", [ext] ExtMethodResponse, }); @@ -84,11 +80,7 @@ impl_jsonrpc_request_enum!(AgentRequest { KillTerminalRequest => "terminal/kill", CreateElicitationRequest => "elicitation/create", #[cfg(feature = "unstable_mcp_over_acp")] - ConnectMcpRequest => "mcp/connect", - #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpRequest => "mcp/message", - #[cfg(feature = "unstable_mcp_over_acp")] - DisconnectMcpRequest => "mcp/disconnect", [ext] ExtMethodRequest, }); @@ -103,18 +95,12 @@ impl_jsonrpc_response_enum!(ClientResponse { KillTerminalResponse => "terminal/kill", CreateElicitationResponse => "elicitation/create", #[cfg(feature = "unstable_mcp_over_acp")] - ConnectMcpResponse => "mcp/connect", - #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpResponse => "mcp/message", - #[cfg(feature = "unstable_mcp_over_acp")] - DisconnectMcpResponse => "mcp/disconnect", [ext] ExtMethodResponse, }); impl_jsonrpc_notification_enum!(AgentNotification { SessionNotification => "session/update", CompleteElicitationNotification => "elicitation/complete", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpNotification => "mcp/message", [ext] ExtNotification, }); diff --git a/src/agent-client-protocol/src/schema/mcp.rs b/src/agent-client-protocol/src/schema/mcp.rs index a16464bc..e5c5032d 100644 --- a/src/agent-client-protocol/src/schema/mcp.rs +++ b/src/agent-client-protocol/src/schema/mcp.rs @@ -1,15 +1,6 @@ //! JSON-RPC implementations for the unstable native MCP-over-ACP transport. -use crate::schema::v1::{ - ConnectMcpRequest, ConnectMcpResponse, DisconnectMcpRequest, DisconnectMcpResponse, - MessageMcpNotification, MessageMcpRequest, MessageMcpResponse, -}; +use crate::schema::v1::{MessageMcpNotification, MessageMcpRequest, MessageMcpResponse}; -impl_jsonrpc_request!(ConnectMcpRequest, ConnectMcpResponse, "mcp/connect"); impl_jsonrpc_request!(MessageMcpRequest, MessageMcpResponse, "mcp/message"); impl_jsonrpc_notification!(MessageMcpNotification, "mcp/message"); -impl_jsonrpc_request!( - DisconnectMcpRequest, - DisconnectMcpResponse, - "mcp/disconnect" -); diff --git a/src/agent-client-protocol/src/schema/v2_impls.rs b/src/agent-client-protocol/src/schema/v2_impls.rs index 8cae486b..94d849e8 100644 --- a/src/agent-client-protocol/src/schema/v2_impls.rs +++ b/src/agent-client-protocol/src/schema/v2_impls.rs @@ -281,14 +281,6 @@ impl_v2_jsonrpc_request!( v2::CreateElicitationResponse, "elicitation/create" ); -#[cfg(feature = "unstable_mcp_over_acp")] -impl_v2_jsonrpc_request!(v2::ConnectMcpRequest, v2::ConnectMcpResponse, "mcp/connect"); -#[cfg(feature = "unstable_mcp_over_acp")] -impl_v2_jsonrpc_request!( - v2::DisconnectMcpRequest, - v2::DisconnectMcpResponse, - "mcp/disconnect" -); impl_v2_jsonrpc_notification!(v2::UpdateSessionNotification, "session/update"); impl_v2_jsonrpc_notification!(v2::CompleteElicitationNotification, "elicitation/complete"); @@ -316,8 +308,6 @@ impl_v2_jsonrpc_request_enum!(v2::ClientRequest { CloseSessionRequest => "session/close", SetSessionConfigOptionRequest => "session/set_config_option", PromptRequest => "session/prompt", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpRequest => "mcp/message", [ext] ExtMethodRequest, }); @@ -340,8 +330,6 @@ impl_v2_jsonrpc_response_enum!(v2::AgentResponse { CloseSessionResponse => "session/close", SetSessionConfigOptionResponse => "session/set_config_option", PromptResponse => "session/prompt", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpResponse => "mcp/message", [ext] ExtMethodResponse, }); @@ -356,11 +344,7 @@ impl_v2_jsonrpc_request_enum!(v2::AgentRequest { RequestPermissionRequest => "session/request_permission", CreateElicitationRequest => "elicitation/create", #[cfg(feature = "unstable_mcp_over_acp")] - ConnectMcpRequest => "mcp/connect", - #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpRequest => "mcp/message", - #[cfg(feature = "unstable_mcp_over_acp")] - DisconnectMcpRequest => "mcp/disconnect", [ext] ExtMethodRequest, }); @@ -368,18 +352,12 @@ impl_v2_jsonrpc_response_enum!(v2::ClientResponse { RequestPermissionResponse => "session/request_permission", CreateElicitationResponse => "elicitation/create", #[cfg(feature = "unstable_mcp_over_acp")] - ConnectMcpResponse => "mcp/connect", - #[cfg(feature = "unstable_mcp_over_acp")] MessageMcpResponse => "mcp/message", - #[cfg(feature = "unstable_mcp_over_acp")] - DisconnectMcpResponse => "mcp/disconnect", [ext] ExtMethodResponse, }); impl_v2_jsonrpc_notification_enum!(v2::AgentNotification { UpdateSessionNotification => "session/update", CompleteElicitationNotification => "elicitation/complete", - #[cfg(feature = "unstable_mcp_over_acp")] - MessageMcpNotification => "mcp/message", [ext] ExtNotification, }); diff --git a/src/agent-client-protocol/src/util.rs b/src/agent-client-protocol/src/util.rs index ed704bcb..dfc36ef9 100644 --- a/src/agent-client-protocol/src/util.rs +++ b/src/agent-client-protocol/src/util.rs @@ -1,10 +1,5 @@ // Types re-exported from crate root -use futures::{ - future::BoxFuture, - stream::{Stream, StreamExt}, -}; - mod typed; pub use typed::{MatchDispatch, MatchDispatchFrom, TypeNotification}; @@ -113,72 +108,3 @@ pub fn run_until( } }) } - -/// Process items from a stream concurrently. -/// -/// For each item received from `stream`, calls `process_fn` to create a future, -/// then runs all futures concurrently. If any future returns an error, -/// stops processing and returns that error. -/// -/// This is useful for patterns where you receive work items from a channel -/// and want to process them concurrently while respecting backpressure. -pub(crate) async fn process_stream_concurrently( - stream: impl Stream, - process_fn: F, - process_fn_hack: impl for<'a> Fn(&'a F, T) -> BoxFuture<'a, Result<(), crate::Error>>, -) -> Result<(), crate::Error> -where - F: AsyncFn(T) -> Result<(), crate::Error>, -{ - use std::pin::pin; - - use futures::stream::{FusedStream, FuturesUnordered}; - use futures_concurrency::future::Race; - - enum Event { - NewItem(Option), - FutureCompleted(Option>), - } - - let mut stream = pin!(stream.fuse()); - let mut futures: FuturesUnordered<_> = FuturesUnordered::new(); - - loop { - // If we have no futures to run, wait until we do. - if futures.is_empty() { - match stream.next().await { - Some(item) => futures.push(process_fn_hack(&process_fn, item)), - None => return Ok(()), - } - continue; - } - - // If there are no more items coming in, just drain our queue and return. - if stream.is_terminated() { - while let Some(result) = futures.next().await { - result?; - } - return Ok(()); - } - - // Otherwise, race between getting a new item and completing a future. - let event = (async { Event::NewItem(stream.next().await) }, async { - Event::FutureCompleted(futures.next().await) - }) - .race() - .await; - - match event { - Event::NewItem(Some(item)) => { - futures.push(process_fn_hack(&process_fn, item)); - } - Event::FutureCompleted(Some(result)) => { - result?; - } - Event::NewItem(None) | Event::FutureCompleted(None) => { - // Stream closed, loop will catch is_terminated - // No futures were pending, shouldn't happen since we checked is_empty - } - } - } -} diff --git a/src/agent-client-protocol/tests/application_dispatch_v2.rs b/src/agent-client-protocol/tests/application_dispatch_v2.rs index 18943e1f..31267f4d 100644 --- a/src/agent-client-protocol/tests/application_dispatch_v2.rs +++ b/src/agent-client-protocol/tests/application_dispatch_v2.rs @@ -3,8 +3,8 @@ use std::{cell::RefCell, rc::Rc, time::Duration}; use agent_client_protocol::{ - Agent, Channel, Client, Error, RawJsonRpcMessage, TransportBatch, TransportFrame, - V2ConnectionTo, + Agent, BudgetedFrame, Channel, Client, Error, RawJsonRpcMessage, TransportBatch, + TransportFrame, V2ConnectionTo, schema::{ProtocolVersion, v2}, }; use futures::{StreamExt as _, channel::mpsc}; @@ -119,13 +119,13 @@ async fn assert_application_order(batched: bool) { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(initialize))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected initialize"); }; assert_eq!(initialize.method.as_ref(), "initialize"); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( initialize.id, Ok(serde_json::to_value( v2::InitializeResponse::new( @@ -138,9 +138,11 @@ async fn assert_application_order(batched: bool) { ) .unwrap()), ))) + .await .unwrap(); - let Some(TransportFrame::Single(RawJsonRpcMessage::Request(resume))) = peer.rx.next().await + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(resume))) = + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected resume"); }; @@ -168,14 +170,16 @@ async fn assert_application_order(batched: bool) { ]; if batched { peer.tx - .unbounded_send(TransportFrame::Batch( + .send_frame(TransportFrame::Batch( TransportBatch::from_messages(messages).unwrap(), )) + .await .unwrap(); } else { for message in messages { peer.tx - .unbounded_send(TransportFrame::Single(message)) + .send_frame(TransportFrame::Single(message)) + .await .unwrap(); } } diff --git a/src/agent-client-protocol/tests/jsonrpc_advanced.rs b/src/agent-client-protocol/tests/jsonrpc_advanced.rs index 1ed60d53..4f69f6e6 100644 --- a/src/agent-client-protocol/tests/jsonrpc_advanced.rs +++ b/src/agent-client-protocol/tests/jsonrpc_advanced.rs @@ -6,7 +6,7 @@ //! - Out-of-order response handling use agent_client_protocol::{ - Channel, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, JsonRpcMessage, + BudgetedFrame, Channel, ConnectionTo, Dispatch, HandleDispatchFrom, Handled, JsonRpcMessage, JsonRpcNotification, JsonRpcRequest, JsonRpcResponse, RawJsonRpcMessage, Responder, SentRequest, TransportBatch, TransportFrame, role::UntypedRole, }; @@ -461,7 +461,7 @@ async fn ordered_callback_installs_dynamic_handler_before_later_batch_entry() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a ping request"); }; @@ -481,7 +481,8 @@ async fn ordered_callback_installs_dynamic_handler_before_later_batch_entry() { let batch = TransportBatch::from_messages([response, notification]) .expect("test batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept the response batch"); while peer.rx.next().await.is_some() {} diff --git a/src/agent-client-protocol/tests/jsonrpc_batch.rs b/src/agent-client-protocol/tests/jsonrpc_batch.rs index 97df192e..8477e3ae 100644 --- a/src/agent-client-protocol/tests/jsonrpc_batch.rs +++ b/src/agent-client-protocol/tests/jsonrpc_batch.rs @@ -865,14 +865,15 @@ async fn protocol_actor_ignores_response_shaped_malformed_public_frame_entries() ]) .expect("test batch is non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("server should accept the test batch"); let frame = tokio::time::timeout(TIMEOUT, peer.rx.next()) .await .expect("timed out waiting for the batch response") .expect("server channel closed before responding"); - let TransportFrame::Batch(batch) = frame else { + let TransportFrame::Batch(batch) = frame.into_frame() else { panic!("request sibling should receive one grouped batch response"); }; let response = serde_json::to_value(batch).expect("batch response should serialize"); diff --git a/src/agent-client-protocol/tests/jsonrpc_error_handling.rs b/src/agent-client-protocol/tests/jsonrpc_error_handling.rs index b1376a36..01fb11e2 100644 --- a/src/agent-client-protocol/tests/jsonrpc_error_handling.rs +++ b/src/agent-client-protocol/tests/jsonrpc_error_handling.rs @@ -90,14 +90,17 @@ async fn response_dispatch_handler_error_reaches_the_local_request_awaiter() { .next() .await .expect("connection should send one request"); - let TransportFrame::Single(RawJsonRpcMessage::Request(request)) = frame else { + let TransportFrame::Single(RawJsonRpcMessage::Request(request)) = + frame.into_frame() + else { panic!("expected one standalone request"); }; peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Ok(serde_json::json!({ "result": "ignored" })), ))) + .await .expect("connection should accept the test response"); Ok::<(), agent_client_protocol::Error>(()) }; diff --git a/src/agent-client-protocol/tests/jsonrpc_request_cancellation.rs b/src/agent-client-protocol/tests/jsonrpc_request_cancellation.rs index 3be9a7c2..dd9fe176 100644 --- a/src/agent-client-protocol/tests/jsonrpc_request_cancellation.rs +++ b/src/agent-client-protocol/tests/jsonrpc_request_cancellation.rs @@ -2,11 +2,10 @@ //! //! These tests avoid sleeps by relying on two ordering guarantees: //! -//! - Messages are delivered in the order they were sent, and each side's -//! dispatch loop processes incoming messages sequentially. A request/response -//! round trip therefore acts as a barrier: by the time the response arrives, -//! every message sent before the request (including any `$/cancel_request`) -//! has been fully processed by the peer. +//! - Ordinary messages preserve queue order, and incoming dispatch is +//! sequential. A round trip after a request proves it reached the peer. +//! Cancellation has a separate urgent lane: canceling before publication +//! settles locally instead of sending either message to the peer. //! - Test handlers report observed cancellations through in-process channels, //! which the test awaits (with a timeout) instead of sleeping. @@ -53,6 +52,22 @@ async fn next_with_timeout(rx: &mut mpsc::UnboundedReceiver) -> T { .expect("channel closed before expected event") } +/// Remote-cancellation tests must first publish the request through every hop. +async fn publication_barrier(connection: &ConnectionTo) { + let response = tokio::time::timeout( + tokio::time::Duration::from_secs(10), + connection + .send_request(SimpleRequest { + message: "barrier".into(), + }) + .block_task(), + ) + .await + .expect("publication barrier timed out") + .expect("publication barrier failed"); + assert_eq!(response.result, "echo: barrier"); +} + /// Assert that no item is currently buffered on `rx`. /// /// Callers must first establish an ordering barrier (such as a @@ -432,6 +447,7 @@ async fn cancelling_request_sent_to_successor_peer_sends_wrapped_cancel() { .run_until(async { let (wrapped_cancel_tx, mut wrapped_cancel_rx) = mpsc::unbounded(); let (plain_cancel_tx, mut plain_cancel_rx) = mpsc::unbounded(); + let (started_tx, mut started_rx) = mpsc::unbounded(); let (server_reader, server_writer, client_reader, client_writer) = setup_test_streams(); let server_transport = @@ -440,9 +456,10 @@ async fn cancelling_request_sent_to_successor_peer_sends_wrapped_cancel() { .builder() .on_receive_request_from( WrappedSuccessor, - async |_request: SimpleRequest, - responder: Responder, - cx: ConnectionTo| { + async move |_request: SimpleRequest, + responder: Responder, + cx: ConnectionTo| { + started_tx.unbounded_send(responder.id().clone()).unwrap(); let cancellation = responder.cancellation(); cx.spawn(async move { let response = cancellation @@ -498,6 +515,7 @@ async fn cancelling_request_sent_to_successor_peer_sends_wrapped_cancel() { }, ); let expected_id = request.id().clone(); + assert_eq!(next_with_timeout(&mut started_rx).await, expected_id); request.cancel()?; let error = request .block_task() @@ -573,6 +591,7 @@ async fn sent_request_can_send_cancellation_for_its_id() { local .run_until(async { let (cancel_tx, mut cancel_rx) = mpsc::unbounded(); + let (started_tx, mut started_rx) = mpsc::unbounded(); let (server_reader, server_writer, client_reader, client_writer) = setup_test_streams(); let server_transport = @@ -580,9 +599,9 @@ async fn sent_request_can_send_cancellation_for_its_id() { let server = UntypedRole .builder() .on_receive_request( - async |request: SimpleRequest, - responder: Responder, - _connection: ConnectionTo| { + async move |request: SimpleRequest, + responder: Responder, + _connection: ConnectionTo| { if request.message == "barrier" { return responder.respond(SimpleResponse { result: format!("echo: {}", request.message), @@ -591,6 +610,7 @@ async fn sent_request_can_send_cancellation_for_its_id() { // Park other requests (by dropping the responder) so // the cancelled request is never answered and the // client handle stays unconsumed. + started_tx.unbounded_send(responder.id().clone()).unwrap(); Ok(()) }, agent_client_protocol::on_receive_request!(), @@ -619,6 +639,7 @@ async fn sent_request_can_send_cancellation_for_its_id() { message: "slow".into(), }); let expected_id = request.id().clone(); + assert_eq!(next_with_timeout(&mut started_rx).await, expected_id); request.cancel()?; let received = next_with_timeout(&mut cancel_rx).await; @@ -1367,6 +1388,7 @@ async fn forward_response_to_propagates_cancellation_to_downstream_request() { connection.send_request(SimpleRequest { message: "cancel downstream".into(), }); + publication_barrier(&connection).await; request.cancel()?; // The backend answers the parked request only once the @@ -1521,6 +1543,7 @@ async fn send_proxied_message_does_not_tunnel_cancel_notifications() { message: "park".into(), }); let client_request_id = request.id().clone(); + publication_barrier(&connection).await; request.cancel()?; let error = request @@ -1824,6 +1847,7 @@ async fn custom_forwarding_propagates_cancellation_when_opted_in() { message: "park".into(), }); let client_request_id = request.id().clone(); + publication_barrier(&connection).await; request.cancel()?; let error = request @@ -1905,6 +1929,7 @@ async fn custom_forwarding_absorbs_cancellation_by_default() { connection.send_request(SimpleRequest { message: "park".into(), }); + publication_barrier(&connection).await; request.cancel()?; // Barrier: the cancellation has now been processed by the diff --git a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs index 4643af1c..7e80ec64 100644 --- a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs +++ b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs @@ -12,14 +12,14 @@ use std::{ }; use agent_client_protocol::{ - ByteStreams, Channel, ConnectTo, ConnectionTo, Dispatch, Error, Handled, JsonRpcMessage, - JsonRpcRequest, Lines, RawJsonRpcMessage, TransportFrame, UntypedMessage, - is_incoming_transport_closed, + BudgetedFrame, ByteStreams, Channel, ConnectTo, ConnectionTo, Dispatch, Error, Handled, + JsonRpcMessage, JsonRpcRequest, Lines, RawJsonRpcMessage, RawJsonRpcResponse as Response, + TransportFrame, UntypedMessage, is_incoming_transport_closed, role::{Role, UntypedRole}, - schema::v1::{RequestId, Response}, + schema::v1::RequestId, }; use agent_client_protocol_test::{MyRequest, MyResponse}; -use futures::{FutureExt as _, SinkExt as _, StreamExt as _, future::join, stream}; +use futures::{FutureExt as _, StreamExt as _, future::join, stream}; use tokio::io::{AsyncBufReadExt as _, AsyncWriteExt as _}; use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _}; @@ -90,8 +90,7 @@ impl ConnectTo for PendingTransport { struct QueuedClient { started: futures::channel::oneshot::Sender<()>, - escaped: - futures::channel::oneshot::Sender>, + escaped: futures::channel::oneshot::Sender, } impl ConnectTo for QueuedClient { @@ -103,7 +102,8 @@ impl ConnectTo for QueuedClient { )?; channel .tx - .unbounded_send(TransportFrame::Single(message)) + .send_frame(TransportFrame::Single(message)) + .await .map_err(Error::into_internal_error)?; drop(self.escaped.send(channel.tx.clone())); let _ = self.started.send(()); @@ -180,7 +180,7 @@ fn assert_connection_closed(error: &Error, method: &str) { async fn receive_requests_then_close(mut peer: Channel, count: usize) { for _ in 0..count { assert!(matches!( - peer.rx.next().await, + peer.rx.next().await.map(BudgetedFrame::into_frame), Some(TransportFrame::Single(RawJsonRpcMessage::Request(_))) )); } @@ -188,12 +188,13 @@ async fn receive_requests_then_close(mut peer: Channel, count: usize) { } async fn respond_then_close(mut peer: Channel) { - let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = peer.rx.next().await + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected outgoing request"); }; peer.tx - .send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Ok(serde_json::to_value(MyResponse { status: "received".into(), @@ -326,7 +327,7 @@ async fn channel_peer_receives_final_response_after_write_half_closes() { let peer = async move { let Channel { mut rx, tx } = peer; - tx.unbounded_send(TransportFrame::Single( + tx.send_frame(TransportFrame::Single( RawJsonRpcMessage::request( "myRequest".into(), serde_json::json!({}), @@ -334,6 +335,7 @@ async fn channel_peer_receives_final_response_after_write_half_closes() { ) .unwrap(), )) + .await .expect("channel should accept the final request"); tx.close_channel(); drop(tx); @@ -341,7 +343,7 @@ async fn channel_peer_receives_final_response_after_write_half_closes() { let Some(TransportFrame::Single(RawJsonRpcMessage::Response(Response::Result { id, result, - }))) = rx.next().await + }))) = rx.next().await.map(BudgetedFrame::into_frame) else { panic!("channel read half closed before the final response"); }; @@ -380,7 +382,7 @@ async fn transport_channel_keeps_read_half_open_after_write_half_closes() { let (channel, transport_future) = ConnectTo::::into_channel_and_future(transport); let Channel { mut rx, tx } = channel; - tx.unbounded_send(TransportFrame::Single( + tx.send_frame(TransportFrame::Single( RawJsonRpcMessage::request( "myRequest".into(), serde_json::json!({}), @@ -388,6 +390,7 @@ async fn transport_channel_keeps_read_half_open_after_write_half_closes() { ) .unwrap(), )) + .await .expect("transport channel should accept the request"); tx.close_channel(); drop(tx); @@ -422,7 +425,7 @@ async fn transport_channel_keeps_read_half_open_after_write_half_closes() { let Some(TransportFrame::Single(RawJsonRpcMessage::Response(Response::Result { id, result, - }))) = rx.next().await + }))) = rx.next().await.map(BudgetedFrame::into_frame) else { panic!("read half closed before delivering the peer's final response"); }; @@ -531,7 +534,7 @@ async fn outgoing_drain_keeps_the_full_duplex_read_half_moving() { assert!( escaped - .unbounded_send(TransportFrame::Single( + .try_send(TransportFrame::Single( RawJsonRpcMessage::notification("too-late".into(), serde_json::json!({}),).unwrap() )) .is_err(), @@ -1077,7 +1080,7 @@ async fn request_finishing_conversion_after_eof_keeps_the_eof_cause() { let connection = tokio::spawn(connection); peer.tx - .send(TransportFrame::Single( + .send_frame(TransportFrame::Single( RawJsonRpcMessage::request( "myRequest".into(), serde_json::json!({}), @@ -1094,7 +1097,8 @@ async fn request_finishing_conversion_after_eof_keeps_the_eof_cause() { assert!(matches!( tokio::time::timeout(TIMEOUT, peer.rx.next()) .await - .expect("handler response was not sent"), + .expect("handler response was not sent") + .map(BudgetedFrame::into_frame), Some(TransportFrame::Single(RawJsonRpcMessage::Response(_))) )); @@ -1233,12 +1237,12 @@ async fn response_buffered_before_eof_is_delivered() { }); let respond_then_close = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected outgoing request"); }; peer.tx - .send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Ok(serde_json::to_value(MyResponse { status: "received".into(), diff --git a/src/agent-client-protocol/tests/live_task_capacity.rs b/src/agent-client-protocol/tests/live_task_capacity.rs new file mode 100644 index 00000000..51b69018 --- /dev/null +++ b/src/agent-client-protocol/tests/live_task_capacity.rs @@ -0,0 +1,80 @@ +use std::time::Duration; + +use agent_client_protocol::{ + BudgetedFrame, Channel, ConnectionLimits, Error, JsonRpcRequest, JsonRpcResponse, + RawJsonRpcMessage, TransportFrame, UntypedRole, +}; +use futures::{StreamExt as _, channel::oneshot}; +use serde::{Deserialize, Serialize}; + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcRequest)] +#[request(method = "_test/callback", response = CallbackResponse)] +struct CallbackRequest {} + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcResponse)] +struct CallbackResponse {} + +#[tokio::test] +async fn persistent_child_cannot_strand_an_ordered_response_callback() -> Result<(), Error> { + tokio::time::timeout(Duration::from_secs(10), async { + let (transport, mut peer) = Channel::duplex_with_limits(ConnectionLimits { + max_queued_frames: 1, + ..ConnectionLimits::default() + }); + let (published_tx, published_rx) = oneshot::channel(); + let (reply_tx, reply_rx) = oneshot::channel::<()>(); + let (child_transport, child_peer) = Channel::duplex(); + let connection = UntypedRole + .builder() + .connect_with(transport, async move |cx| { + let _child = cx.spawn_connection(UntypedRole.builder(), child_transport)?; + let (callback_tx, callback_rx) = oneshot::channel(); + let request = cx.send_request(CallbackRequest {}); + published_rx.await.map_err(Error::into_internal_error)?; + let error = request + .on_receiving_result(async move |_result| { + let _ = callback_tx.send(()); + Ok(()) + }) + .expect_err("the child reserves the sole live task slot"); + assert!(error.to_string().contains("live task capacity"), "{error}"); + assert!( + callback_rx.await.is_err(), + "rejected callback must release its captured resources" + ); + let _ = reply_tx.send(()); + cx.incoming_closed().await; + Ok(()) + }); + let reply_and_close = async move { + let request = loop { + match peer.rx.next().await.map(BudgetedFrame::into_frame) { + Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) => { + break request; + } + Some(TransportFrame::Single(RawJsonRpcMessage::Notification(_))) => {} + other => panic!("parent request did not reach the peer: {other:?}"), + } + }; + let _ = published_tx.send(()); + let _ = reply_rx.await; + peer.tx + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( + request.id, + Ok(serde_json::json!({})), + ))) + .await + .unwrap(); + // Half-close input but continue draining any cancellation/output. + // Dropping both directions here would make a legitimate late write + // fail for reasons unrelated to the response acknowledgment. + peer.tx.close_channel(); + while peer.rx.next().await.is_some() {} + }; + let (connection, ()) = futures::future::join(connection, reply_and_close).await; + drop(child_peer); + connection + }) + .await + .expect("persistent child stranded the response dispatcher") +} diff --git a/src/agent-client-protocol/tests/mcp_connector_admission.rs b/src/agent-client-protocol/tests/mcp_connector_admission.rs new file mode 100644 index 00000000..70155ff6 --- /dev/null +++ b/src/agent-client-protocol/tests/mcp_connector_admission.rs @@ -0,0 +1,336 @@ +#![cfg(feature = "unstable_mcp_over_acp")] + +use std::{ + future::pending, + sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, ByteStreams, Channel, Client, ConnectTo, ConnectionLimits, ConnectionTo, DynConnectTo, + Error, FrameSender, RawJsonRpcMessage, Responder, RunWithConnectionTo, TransportFrame, + mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, + role, + schema::v1, +}; +use futures::StreamExt as _; +use serde_json::{Map, Value, json}; +use tokio::{ + io::{AsyncBufReadExt, AsyncWriteExt, BufReader, duplex}, + sync::oneshot, +}; +use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; + +const TIMEOUT: Duration = Duration::from_secs(10); + +#[derive(Default)] +struct Probes { + factory: AtomicUsize, + backend_started: AtomicUsize, + backend_dropped: AtomicUsize, + escaped_senders: Mutex>, + notifications: Mutex>, +} + +#[derive(Clone, Copy)] +enum BackendBehavior { + WireReply, + ReplyThenError, + ExitWithoutReply, +} + +struct Connector(Arc, BackendBehavior); + +impl McpServerConnect for Connector { + fn name(&self) -> String { + "capacity-probe".into() + } + + fn connect(&self, context: McpConnectionTo) -> DynConnectTo { + assert!(context.request_id().is_some()); + self.0.factory.fetch_add(1, Ordering::SeqCst); + DynConnectTo::new(Backend(self.0.clone(), self.1)) + } +} + +struct Backend(Arc, BackendBehavior); + +struct BackendDrop(Arc); + +impl Drop for BackendDrop { + fn drop(&mut self) { + self.0.backend_dropped.fetch_add(1, Ordering::SeqCst); + } +} + +impl ConnectTo for Backend { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + self.0.backend_started.fetch_add(1, Ordering::SeqCst); + let _drop = BackendDrop(self.0.clone()); + if !matches!(self.1, BackendBehavior::WireReply) { + let (mut channel, driver) = client.into_channel_and_future(); + let work = async { + let frame = channel.rx.next().await.expect("MCP request"); + let TransportFrame::Single(RawJsonRpcMessage::Request(request)) = + frame.into_frame() + else { + panic!("expected one request"); + }; + self.0 + .escaped_senders + .lock() + .unwrap() + .push(channel.tx.clone()); + if matches!(self.1, BackendBehavior::ExitWithoutReply) { + return Ok(()); + } + for message in [ + RawJsonRpcMessage::notification( + "notifications/progress".into(), + json!({"progressToken":1, "progress":1, "marker":"before"}), + )?, + RawJsonRpcMessage::response(request.id, Ok(json!({"admitted":true}))), + RawJsonRpcMessage::notification( + "notifications/progress".into(), + json!({"progressToken":1, "progress":2, "marker":"after"}), + )?, + ] { + channel + .tx + .try_send(TransportFrame::Single(message)) + .map_err(Error::into_internal_error)?; + } + Err(Error::internal_error().data("driver failed after accepted output")) + }; + let (driver, result) = tokio::join!(driver, work); + driver?; + return result; + } + let (sdk_output, peer_input) = duplex(4096); + let (peer_output, sdk_input) = duplex(4096); + let transport = ByteStreams::new(sdk_output.compat_write(), sdk_input.compat()); + let peer = async move { + let mut reader = BufReader::new(peer_input); + let mut line = String::new(); + if reader.read_line(&mut line).await.expect("read MCP request") == 0 { + // Rejected admissions close the backend before sending anything. + return; + } + let request: Value = serde_json::from_str(&line).expect("valid MCP request"); + assert_eq!( + request["params"]["_meta"]["io.modelcontextprotocol/protocolVersion"], + "2026-07-28" + ); + let response = json!({ + "jsonrpc": "2.0", + "id": request["id"], + "result": {"admitted": true} + }); + let mut output = peer_output; + output + .write_all(format!("{response}\n").as_bytes()) + .await + .expect("write MCP response"); + output.shutdown().await.expect("close MCP backend output"); + }; + let (result, ()) = tokio::join!(client.connect_to(transport), peer); + result + } +} + +struct NullRun; + +impl RunWithConnectionTo for NullRun { + async fn run_with_connection_to(self, _connection: ConnectionTo) -> Result<(), Error> { + pending().await + } +} + +// Unlike a raw Channel, this uses ConnectTo's default adapter. Only the +// provider endpoint is limited; the agent must have its own default pool. +struct DefaultCapacityAgent(Channel); + +impl ConnectTo for DefaultCapacityAgent { + async fn connect_to(self, agent: impl ConnectTo) -> Result<(), Error> { + ConnectTo::::connect_to(self.0, agent).await + } +} + +fn params() -> Map { + serde_json::from_value(json!({ + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {}, + "progressToken": 1 + } + })) + .unwrap() +} + +// Occupy slots only after newSession has completed. The agent waits on `start` +// so no MCP request can race the provider's admission setup. +async fn scenario( + fillers: usize, + requests: usize, + behavior: BackendBehavior, +) -> ( + Vec<(Result, usize)>, + Arc, +) { + tokio::time::timeout(TIMEOUT, async move { + let probes = Arc::new(Probes::default()); + let (provider_channel, agent_channel) = Channel::duplex_with_limits(ConnectionLimits { + max_queued_frames: 8, + ..ConnectionLimits::default() + }); + let (start_tx, start_rx) = oneshot::channel::<()>(); + let start_rx = Mutex::new(Some(start_rx)); + let (done_tx, done_rx) = oneshot::channel(); + let done_tx = Mutex::new(Some(done_tx)); + let agent_probes = probes.clone(); + let notification_probes = probes.clone(); + let agent = Agent.builder().on_receive_request( + async move |request: v1::NewSessionRequest, + responder: Responder, + connection: ConnectionTo| { + let [v1::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected one native MCP server") + }; + let server_id = server.server_id.clone(); + responder.respond(v1::NewSessionResponse::new(v1::SessionId::new( + "capacity-session", + )))?; + let start = start_rx.lock().unwrap().take().expect("one session"); + let done = done_tx.lock().unwrap().take().expect("one session"); + let sender = connection.clone(); + let observed = agent_probes.clone(); + connection.spawn(async move { + start.await.map_err(Error::into_internal_error)?; + let mut responses = Vec::new(); + for _ in 0..requests { + // Reuse the logical ID after the first request completes. + let response = sender + .send_request( + v1::MessageMcpRequest::new( + server_id.clone(), + "reused-id", + "admission/probe", + ) + .params(params()), + ) + .block_task() + .await; + responses.push((response, observed.backend_dropped.load(Ordering::SeqCst))); + } + drop(done.send(responses)); + Ok(()) + }) + }, + agent_client_protocol::on_receive_request!(), + ); + let agent = agent.on_receive_notification( + async move |notification: v1::MessageMcpNotification, _cx| { + assert_eq!(notification.request_id, v1::McpRequestId::new("reused-id")); + notification_probes.notifications.lock().unwrap().push( + notification.params.as_ref().unwrap()["marker"] + .as_str() + .unwrap() + .to_owned(), + ); + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ); + let connector = Connector(probes.clone(), behavior); + let test = Client + .builder() + .connect_with(provider_channel, async move |connection| { + let filler_connection = connection.clone(); + connection + .build_session_cwd()? + .with_mcp_server(McpServer::::new(connector, NullRun))? + .block_task() + .run_until(async |_session| { + for _ in 0..fillers { + filler_connection + .spawn(async { pending::>().await })?; + } + start_tx.send(()).expect("agent still waiting"); + let responses = done_rx.await.map_err(Error::into_internal_error)?; + Ok(responses) + }) + .await + }); + let (result, agent_result) = + tokio::join!(test, DefaultCapacityAgent(agent_channel).connect_to(agent)); + agent_result.expect("agent connection"); + (result.expect("provider connection"), probes) + }) + .await + .expect("MCP capacity scenario timed out") +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn one_free_slot_completes_connector_and_recovers_for_reused_id() { + let (responses, probes) = scenario(7, 2, BackendBehavior::WireReply).await; + for (index, (response, dropped_at_response)) in responses.into_iter().enumerate() { + match response.expect("one slot must admit the MCP request") { + v1::MessageMcpResponse::Result { result, .. } => { + assert_eq!(result, json!({"admitted": true})); + } + other => panic!("expected MCP result, got {other:?}"), + } + assert_eq!( + dropped_at_response, + index + 1, + "backend must stop before the logical response is observed" + ); + } + assert_eq!(probes.factory.load(Ordering::SeqCst), 2); + assert_eq!(probes.backend_started.load(Ordering::SeqCst), 2); + assert_eq!(probes.backend_dropped.load(Ordering::SeqCst), 2); +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn zero_free_slots_rejects_before_factory_or_backend_setup() { + let (responses, probes) = scenario(8, 1, BackendBehavior::WireReply).await; + let [(response, _)] = <[_; 1]>::try_from(responses).expect("one request"); + let error = response.expect_err("all live slots are occupied"); + assert!(error.to_string().contains("live task capacity"), "{error}"); + assert_eq!(probes.factory.load(Ordering::SeqCst), 0); + assert_eq!(probes.backend_started.load(Ordering::SeqCst), 0); + assert_eq!(probes.backend_dropped.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn connector_drains_terminal_output_before_driver_failure_and_rejects_late_output() { + let (responses, probes) = scenario(7, 1, BackendBehavior::ReplyThenError).await; + let [(response, dropped)] = <[_; 1]>::try_from(responses).unwrap(); + let v1::MessageMcpResponse::Result { result, .. } = response.unwrap() else { + panic!("accepted terminal result must survive backend exit"); + }; + assert_eq!(result, json!({"admitted":true})); + assert_eq!(dropped, 1); + assert_eq!(*probes.notifications.lock().unwrap(), ["before"]); + let escaped = probes.escaped_senders.lock().unwrap(); + assert_eq!(escaped.len(), 1); + assert!(escaped[0].is_closed()); +} + +#[tokio::test] +async fn connector_exit_without_response_does_not_wait_for_escaped_sender() { + let (responses, probes) = scenario(7, 1, BackendBehavior::ExitWithoutReply).await; + let [(response, dropped)] = <[_; 1]>::try_from(responses).unwrap(); + let error = response.expect_err("completed backend did not produce a terminal outcome"); + assert_eq!( + i32::from(error.code), + agent_client_protocol::mcp_server::MCP_BACKEND_FAILURE + ); + assert_eq!(dropped, 1); + let escaped = probes.escaped_senders.lock().unwrap(); + assert_eq!(escaped.len(), 1); + assert!(escaped[0].is_closed()); +} diff --git a/src/agent-client-protocol/tests/mcp_connector_errors.rs b/src/agent-client-protocol/tests/mcp_connector_errors.rs new file mode 100644 index 00000000..6865f44b --- /dev/null +++ b/src/agent-client-protocol/tests/mcp_connector_errors.rs @@ -0,0 +1,313 @@ +#![cfg(feature = "unstable_mcp_over_acp")] + +use std::time::Duration; + +#[cfg(feature = "unstable_protocol_v2")] +use agent_client_protocol::V2ConnectionTo; +#[cfg(feature = "unstable_protocol_v2")] +use agent_client_protocol::schema::v2; +use agent_client_protocol::{ + Agent, ByteStreams, Client, ConnectTo, ConnectionTo, DynConnectTo, Error, Responder, + RunWithConnectionTo, + mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, + role, + schema::v1, +}; +use serde_json::{Map, Value, json}; +use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader, duplex}; +use tokio_util::compat::{TokioAsyncReadCompatExt, TokioAsyncWriteCompatExt}; + +const TIMEOUT: Duration = Duration::from_secs(10); +const CASES: &[(&str, &str)] = &[ + ( + "absent", + r#"{"code":-32000,"message":"backend failed","extension":{"retry":false}}"#, + ), + ( + "null", + r#"{"code":-32000,"message":"backend failed","data":null,"extension":{"retry":false}}"#, + ), + ( + "object", + r#"{"code":-32000,"message":"backend failed","data":{"cause":"upstream"},"extension":{"retry":false}}"#, + ), +]; + +struct WireConnector; + +impl McpServerConnect for WireConnector { + fn name(&self) -> String { + "raw-wire-errors".into() + } + + fn connect(&self, context: McpConnectionTo) -> DynConnectTo { + assert!( + context.request_id().is_some(), + "expected an ACP MCP request" + ); + DynConnectTo::new(WireBackend) + } +} + +struct WireBackend; + +impl ConnectTo for WireBackend { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + let (sdk_output, peer_input) = duplex(4096); + let (peer_output, sdk_input) = duplex(4096); + let transport = ByteStreams::new(sdk_output.compat_write(), sdk_input.compat()); + let peer = async move { + let mut reader = BufReader::new(peer_input); + let mut line = String::new(); + reader + .read_line(&mut line) + .await + .expect("read MCP wire request"); + let request: Value = serde_json::from_str(&line).expect("valid MCP wire request"); + assert_eq!(request["jsonrpc"], "2.0"); + assert_eq!( + request["params"]["_meta"]["io.modelcontextprotocol/protocolVersion"], + "2026-07-28" + ); + let id = serde_json::to_string(&request["id"]).unwrap(); + let method = request["method"].as_str().expect("MCP method"); + let response = if method == "success" { + assert_eq!( + request["id"], "next", + "backend must see the logical request ID" + ); + format!(r#"{{"jsonrpc":"2.0","id":{id},"result":null}}"#) + } else { + let index = CASES + .iter() + .position(|(case, _)| case == &method) + .expect("known test case"); + assert_eq!( + request["id"], + format!("wire-{index}"), + "backend must see the logical request ID" + ); + let (_, error) = CASES + .iter() + .find(|(case, _)| case == &method) + .expect("known test case"); + // Literal backend JSON, not ACP Error or RawJsonRpcMessage::response: + // those typed constructors cannot represent data:null or extension. + format!(r#"{{"jsonrpc":"2.0","id":{id},"error":{error}}}"#) + }; + let mut output = peer_output; + output + .write_all(response.as_bytes()) + .await + .expect("write MCP response"); + output.write_all(b"\n").await.expect("frame MCP response"); + output.shutdown().await.expect("close MCP response stream"); + }; + let (result, ()) = tokio::join!(client.connect_to(transport), peer); + result + } +} + +struct IdleRunner; + +impl RunWithConnectionTo for IdleRunner { + async fn run_with_connection_to(self, _connection: ConnectionTo) -> Result<(), Error> { + std::future::pending().await + } +} + +#[cfg(feature = "unstable_protocol_v2")] +fn cwd() -> std::path::PathBuf { + std::env::current_dir().expect("cwd") +} + +fn params() -> Map { + serde_json::from_value(json!({ + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } + })) + .unwrap() +} + +fn assert_error(error: impl serde::Serialize, case: &str) { + let value = serde_json::to_value(error).unwrap(); + assert_eq!(value["code"], -32000, "{case}"); + assert_eq!(value["message"], "backend failed", "{case}"); + assert_eq!(value["extension"], json!({"retry": false}), "{case}"); + let object = value.as_object().unwrap(); + match case { + "absent" => assert!(!object.contains_key("data"), "{value}"), + "null" => assert_eq!(object.get("data"), Some(&Value::Null)), + "object" => assert_eq!(object.get("data"), Some(&json!({"cause":"upstream"}))), + _ => unreachable!(), + } +} + +async fn v1_requests( + connection: ConnectionTo, + server_id: v1::McpServerAcpId, +) -> Result<(), Error> { + for (index, (case, _)) in CASES.iter().enumerate() { + let response = connection + .send_request( + v1::MessageMcpRequest::new(server_id.clone(), format!("wire-{index}"), *case) + .params(params()), + ) + .block_task() + .await?; + match response { + v1::MessageMcpResponse::Error { error, .. } => assert_error(error, case), + other => panic!("{case}: expected inner MCP error, got {other:?}"), + } + } + let response = connection + .send_request(v1::MessageMcpRequest::new(server_id, "next", "success").params(params())) + .block_task() + .await?; + assert!(matches!( + response, + v1::MessageMcpResponse::Result { + result: Value::Null, + .. + } + )); + Ok(()) +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn v1_connector_preserves_backend_wire_errors() { + let (done_tx, done_rx) = tokio::sync::oneshot::channel(); + let done_tx = std::sync::Mutex::new(Some(done_tx)); + let agent = Agent.builder().on_receive_request( + async move |request: v1::NewSessionRequest, + responder: Responder, + connection: ConnectionTo| { + let [v1::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected one native MCP server") + }; + let server_id = server.server_id.clone(); + responder.respond(v1::NewSessionResponse::new(v1::SessionId::new( + "wire-session", + )))?; + let done = done_tx.lock().unwrap().take().expect("one setup request"); + let requests = connection.clone(); + connection.spawn(async move { + drop(done.send(v1_requests(requests, server_id).await)); + Ok(()) + }) + }, + agent_client_protocol::on_receive_request!(), + ); + let test = Client.builder().connect_with(agent, async |connection| { + connection + .build_session_cwd()? + .with_mcp_server(McpServer::::new(WireConnector, IdleRunner))? + .block_task() + .run_until(async |_session| { + done_rx.await.map_err(Error::into_internal_error)??; + Ok(()) + }) + .await?; + Ok(()) + }); + tokio::time::timeout(TIMEOUT, test) + .await + .expect("v1 connector timed out") + .expect("v1 connector failed"); +} + +#[cfg(feature = "unstable_protocol_v2")] +async fn v2_requests( + connection: V2ConnectionTo, + server_id: v2::McpServerAcpId, +) -> Result<(), Error> { + for (index, (case, _)) in CASES.iter().enumerate() { + let response = connection + .send_request( + v2::MessageMcpRequest::new(server_id.clone(), format!("wire-{index}"), *case) + .params(params()), + ) + .block_task() + .await?; + match response { + v2::MessageMcpResponse::Error { error, .. } => assert_error(error, case), + other => panic!("{case}: expected inner MCP error, got {other:?}"), + } + } + let response = connection + .send_request(v2::MessageMcpRequest::new(server_id, "next", "success").params(params())) + .block_task() + .await?; + assert!(matches!( + response, + v2::MessageMcpResponse::Result { + result: Value::Null, + .. + } + )); + Ok(()) +} + +#[cfg(feature = "unstable_protocol_v2")] +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +async fn v2_connector_preserves_backend_wire_errors() { + let (done_tx, done_rx) = tokio::sync::oneshot::channel(); + let done_tx = std::sync::Mutex::new(Some(done_tx)); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _connection: V2ConnectionTo| { + responder.respond(v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("wire-backend-test", env!("CARGO_PKG_VERSION")), + )) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async move |request: v2::NewSessionRequest, + responder: Responder, + connection: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected one native MCP server") + }; + let server_id = server.server_id.clone(); + responder.respond(v2::NewSessionResponse::new(v2::SessionId::new( + "wire-session", + )))?; + let done = done_tx.lock().unwrap().take().expect("one setup request"); + let requests = connection.clone(); + connection.spawn(async move { + drop(done.send(v2_requests(requests, server_id).await)); + Ok(()) + }) + }, + agent_client_protocol::on_receive_request!(), + ); + let test = Client.v2().connect_with(agent, async |connection| { + connection + .send_request(v2::InitializeRequest::new( + agent_client_protocol::schema::ProtocolVersion::V2, + v2::Implementation::new("wire-backend-test", env!("CARGO_PKG_VERSION")), + )) + .block_task() + .await?; + let session = connection + .build_session_from(v2::NewSessionRequest::new(cwd())) + .with_mcp_server(McpServer::::new(WireConnector, IdleRunner))? + .start_session() + .block_task() + .await?; + done_rx.await.map_err(Error::into_internal_error)??; + drop(session); + Ok(()) + }); + tokio::time::timeout(TIMEOUT, test) + .await + .expect("v2 connector timed out") + .expect("v2 connector failed"); +} diff --git a/src/agent-client-protocol/tests/mcp_message_deserialization.rs b/src/agent-client-protocol/tests/mcp_message_deserialization.rs new file mode 100644 index 00000000..dd8c3fe0 --- /dev/null +++ b/src/agent-client-protocol/tests/mcp_message_deserialization.rs @@ -0,0 +1,73 @@ +//! A shared method name must not conflate JSON-RPC requests and notifications. +#![cfg(feature = "unstable_mcp_over_acp")] + +use agent_client_protocol::{JsonRpcMessage, RawJsonRpcMessage, schema::v1}; +use serde_json::{Value, json}; + +fn params() -> Value { + // Deliberately identical for both message kinds: dispatch must use the outer + // envelope, not infer a kind from this opaque inner method or requestId. + json!({"serverId":"server", "requestId":"logical-id", "method":"custom/message"}) +} + +#[test] +fn mcp_message_kind_is_selected_by_outer_id() { + let notification = json!({"jsonrpc":"2.0", "method":"mcp/message", "params":params()}); + let parsed: RawJsonRpcMessage = serde_json::from_value(notification.clone()).unwrap(); + let RawJsonRpcMessage::Notification(parsed) = parsed else { + panic!("nested requestId must not turn a notification into a request"); + }; + assert_eq!(parsed.method.as_ref(), "mcp/message"); + assert_eq!(parsed.params.unwrap().into_value(), params()); + + for id in [json!(42), json!("outer-id")] { + let mut request = notification.clone(); + request["id"] = id.clone(); + let parsed: RawJsonRpcMessage = serde_json::from_value(request).unwrap(); + let RawJsonRpcMessage::Request(parsed) = parsed else { + panic!("outer id identifies a request"); + }; + assert_eq!(serde_json::to_value(parsed.id).unwrap(), id); + assert_eq!(parsed.method.as_ref(), "mcp/message"); + assert_eq!(parsed.params.unwrap().into_value(), params()); + } + + for id in [json!(true), json!({}), json!([])] { + let mut malformed = notification.clone(); + malformed["id"] = id; + assert!( + serde_json::from_value::(malformed).is_err(), + "an invalid request id must not fall back to notification deserialization" + ); + } +} + +#[test] +fn v1_mcp_method_is_in_separate_request_and_notification_enums() { + assert!(matches!( + v1::AgentRequest::parse_message("mcp/message", ¶ms()).unwrap(), + v1::AgentRequest::MessageMcpRequest(_) + )); + assert!(matches!( + v1::ClientNotification::parse_message("mcp/message", ¶ms()).unwrap(), + v1::ClientNotification::MessageMcpNotification(_) + )); + assert!(v1::ClientRequest::parse_message("mcp/message", ¶ms()).is_err()); + assert!(v1::AgentNotification::parse_message("mcp/message", ¶ms()).is_err()); +} + +#[cfg(feature = "unstable_protocol_v2")] +#[test] +fn v2_mcp_method_is_in_separate_request_and_notification_enums() { + use agent_client_protocol::schema::v2; + assert!(matches!( + v2::AgentRequest::parse_message("mcp/message", ¶ms()).unwrap(), + v2::AgentRequest::MessageMcpRequest(_) + )); + assert!(matches!( + v2::ClientNotification::parse_message("mcp/message", ¶ms()).unwrap(), + v2::ClientNotification::MessageMcpNotification(_) + )); + assert!(v2::ClientRequest::parse_message("mcp/message", ¶ms()).is_err()); + assert!(v2::AgentNotification::parse_message("mcp/message", ¶ms()).is_err()); +} diff --git a/src/agent-client-protocol/tests/meta_propagation.rs b/src/agent-client-protocol/tests/meta_propagation.rs index eef3f331..537ea09f 100644 --- a/src/agent-client-protocol/tests/meta_propagation.rs +++ b/src/agent-client-protocol/tests/meta_propagation.rs @@ -106,7 +106,7 @@ fn successor_message_accepts_legacy_meta_alias() -> Result<(), agent_client_prot fn native_mcp_over_acp_message_meta_serializes_as_reserved_meta_field() -> Result<(), agent_client_protocol::Error> { let meta = trace_context_meta(); - let message = MessageMcpRequest::new("connection-1", "tools/list") + let message = MessageMcpRequest::new("server-1", "request-1", "tools/list") .params(serde_json::Map::from_iter([( "cursor".into(), Value::String("abc".into()), @@ -116,7 +116,8 @@ fn native_mcp_over_acp_message_meta_serializes_as_reserved_meta_field() let untyped = message.to_untyped_message()?; assert_eq!(untyped.method(), "mcp/message"); - assert_eq!(untyped.params()["connectionId"], "connection-1"); + assert_eq!(untyped.params()["serverId"], "server-1"); + assert_eq!(untyped.params()["requestId"], "request-1"); assert_eq!(untyped.params()["method"], "tools/list"); assert_eq!(untyped.params()["params"]["cursor"], "abc"); assert_eq!(untyped.params()["_meta"], Value::Object(meta.clone())); diff --git a/src/agent-client-protocol/tests/native_mcp_shutdown.rs b/src/agent-client-protocol/tests/native_mcp_shutdown.rs new file mode 100644 index 00000000..8b8fd7e6 --- /dev/null +++ b/src/agent-client-protocol/tests/native_mcp_shutdown.rs @@ -0,0 +1,220 @@ +#![cfg(all(feature = "unstable_protocol_v2", feature = "unstable_mcp_over_acp"))] + +use std::{sync::Mutex, time::Duration}; + +use agent_client_protocol::{ + Agent, Channel, Client, ConnectionTo, Error, Responder, RunWithConnectionTo, V2ConnectionTo, + mcp_server::{McpOutcome, McpRequest, McpRequestContext, McpServer, McpService}, + schema::{ProtocolVersion, v2}, +}; +use futures::future::BoxFuture; +use serde_json::json; +use tokio::sync::oneshot; + +struct OnDrop(Option>); + +impl Drop for OnDrop { + fn drop(&mut self) { + if let Some(done) = self.0.take() { + let _ = done.send(()); + } + } +} + +struct CleanupService { + started: Mutex>>, + cleanup_started: Mutex>>, + runner_woke: Mutex>>, + release: Mutex>>, + dropped: Mutex>>, +} + +struct ShutdownRunner(oneshot::Sender<()>); + +impl RunWithConnectionTo for ShutdownRunner { + async fn run_with_connection_to(self, cx: ConnectionTo) -> Result<(), Error> { + cx.shutdown_requested().await; + let _ = self.0.send(()); + std::future::pending().await + } +} + +impl McpService for CleanupService { + fn execute( + &self, + _request: McpRequest, + context: McpRequestContext, + ) -> BoxFuture<'static, Result> { + let started = self.started.lock().unwrap().take().unwrap(); + let cleanup_started = self.cleanup_started.lock().unwrap().take().unwrap(); + let runner_woke = self.runner_woke.lock().unwrap().take().unwrap(); + let release = self.release.lock().unwrap().take().unwrap(); + let dropped = self.dropped.lock().unwrap().take().unwrap(); + Box::pin(async move { + let _drop = OnDrop(Some(dropped)); + let _ = started.send(()); + context.operation_cancellation().cancelled().await; + runner_woke.await.map_err(Error::into_internal_error)?; + let _ = cleanup_started.send(()); + let _ = release.await; + Err(Error::request_cancelled()) + }) + } +} + +#[derive(Clone, Copy)] +enum Shutdown { + PeerEof, + Foreground, + UnrelatedTaskError, +} + +async fn shutdown_joins_native_cleanup(shutdown: Shutdown) -> Result<(), Error> { + tokio::time::timeout(Duration::from_secs(10), async move { + let (started_tx, started_rx) = oneshot::channel(); + let (cleanup_started_tx, cleanup_started_rx) = oneshot::channel(); + let (runner_woke_tx, runner_woke_rx) = oneshot::channel(); + let (release_tx, release_rx) = oneshot::channel(); + let (dropped_tx, mut dropped_rx) = oneshot::channel(); + let (peer_stop_tx, peer_stop_rx) = oneshot::channel::<()>(); + let mut peer_stop_tx = Some(peer_stop_tx); + let (client_stop_tx, client_stop_rx) = oneshot::channel::<()>(); + let (peer, client) = Channel::duplex(); + let agent = Agent + .v2() + .on_receive_request( + async |request: v2::InitializeRequest, + responder: Responder, + _cx| { + responder.respond(v2::InitializeResponse::new( + request.protocol_version, + v2::Implementation::new("shutdown-agent", "1"), + )) + }, + agent_client_protocol::on_receive_request!(), + ) + .on_receive_request( + async |request: v2::NewSessionRequest, + responder: Responder, + cx: V2ConnectionTo| { + let [v2::McpServer::Acp(server)] = request.mcp_servers.as_slice() else { + panic!("expected native MCP server declaration"); + }; + let server_id = server.server_id.clone(); + let request_connection = cx.clone(); + cx.spawn(async move { + let request = + v2::MessageMcpRequest::new(server_id, "cleanup-probe", "tools/call") + .params( + json!({ + "name": "probe", + "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } + }) + .as_object() + .unwrap() + .clone(), + ); + let _result = request_connection.send_request(request).block_task().await; + Ok(()) + })?; + responder.respond(v2::NewSessionResponse::new("shutdown-session")) + }, + agent_client_protocol::on_receive_request!(), + ); + let peer_task = tokio::spawn(agent.connect_with(peer, async move |_cx| { + let _ = peer_stop_rx.await; + Ok(()) + })); + let client_task = tokio::spawn(Client.v2().connect_with(client, async move |cx| { + cx.send_request(v2::InitializeRequest::new( + ProtocolVersion::V2, + v2::Implementation::new("shutdown-client", "1"), + )) + .block_task() + .await?; + let server = McpServer::new_service( + CleanupService { + started: Mutex::new(Some(started_tx)), + cleanup_started: Mutex::new(Some(cleanup_started_tx)), + runner_woke: Mutex::new(Some(runner_woke_rx)), + release: Mutex::new(Some(release_rx)), + dropped: Mutex::new(Some(dropped_tx)), + }, + "shutdown-test", + ShutdownRunner(runner_woke_tx), + ); + cx.build_session(std::env::current_dir().map_err(Error::into_internal_error)?) + .with_mcp_server(server)? + .start_session() + .block_task() + .await?; + match shutdown { + Shutdown::PeerEof => cx.incoming_closed().await, + Shutdown::Foreground => { + let _ = client_stop_rx.await; + } + Shutdown::UnrelatedTaskError => { + let _ = client_stop_rx.await; + cx.spawn(async { Err(Error::internal_error().data("unrelated task failed")) })?; + std::future::pending::<()>().await; + } + } + Ok(()) + })); + + started_rx.await.map_err(Error::into_internal_error)?; + if matches!(shutdown, Shutdown::PeerEof) { + let _ = peer_stop_tx.take().unwrap().send(()); + } else { + let _ = client_stop_tx.send(()); + } + cleanup_started_rx + .await + .map_err(Error::into_internal_error)?; + assert!( + !client_task.is_finished(), + "connection discarded pending native cleanup" + ); + assert!(matches!( + dropped_rx.try_recv(), + Err(oneshot::error::TryRecvError::Empty) + )); + let _ = release_tx.send(()); + dropped_rx.await.map_err(Error::into_internal_error)?; + let result = client_task.await.map_err(Error::into_internal_error)?; + if matches!(shutdown, Shutdown::UnrelatedTaskError) { + let error = result.expect_err("task failure must remain the primary connection error"); + assert!( + error.to_string().contains("unrelated task failed"), + "{error}" + ); + } else { + result?; + } + if !matches!(shutdown, Shutdown::PeerEof) { + let _ = peer_stop_tx.take().unwrap().send(()); + } + peer_task.await.map_err(Error::into_internal_error)??; + Ok(()) + }) + .await + .expect("native cleanup shutdown timed out") +} + +#[tokio::test] +async fn native_cleanup_survives_clean_incoming_eof() -> Result<(), Error> { + shutdown_joins_native_cleanup(Shutdown::PeerEof).await +} + +#[tokio::test] +async fn native_cleanup_survives_foreground_return() -> Result<(), Error> { + shutdown_joins_native_cleanup(Shutdown::Foreground).await +} + +#[tokio::test] +async fn unrelated_task_error_waits_for_native_cleanup_and_preserves_error() -> Result<(), Error> { + shutdown_joins_native_cleanup(Shutdown::UnrelatedTaskError).await +} diff --git a/src/agent-client-protocol/tests/protocol_v2.rs b/src/agent-client-protocol/tests/protocol_v2.rs index 52a94a4a..8eeab227 100644 --- a/src/agent-client-protocol/tests/protocol_v2.rs +++ b/src/agent-client-protocol/tests/protocol_v2.rs @@ -409,18 +409,23 @@ impl ConnectTo for FutureInitializeV2Client { channel .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "initialize".into(), initialize_params_with_extensions(ProtocolVersion::V2)?, v1::RequestId::Number(1), )?)) + .await .map_err(Error::into_internal_error)?; while let Some(message) = channel.rx.next().await { - let TransportFrame::Single(message) = message else { + let TransportFrame::Single(message) = message.into_frame() else { continue; }; - let RawJsonRpcMessage::Response(v1::Response::Result { result, .. }) = message else { + let RawJsonRpcMessage::Response(agent_client_protocol::RawJsonRpcResponse::Result { + result, + .. + }) = message + else { continue; }; let initialize = v2::InitializeResponse::from_value("initialize", result)?; @@ -488,27 +493,28 @@ async fn assert_malformed_initialize_rejected(params: Map) -> Res channel .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "initialize".into(), Value::Object(params), v1::RequestId::Number(1), )?)) + .await .map_err(Error::into_internal_error)?; while let Some(message) = channel.rx.next().await { - let TransportFrame::Single(message) = message else { + let TransportFrame::Single(message) = message.into_frame() else { continue; }; let RawJsonRpcMessage::Response(response) = message else { continue; }; - let v1::Response::Error { error, .. } = response else { + let agent_client_protocol::RawJsonRpcResponse::Error { error, .. } = response else { panic!("malformed initialize should fail"); }; - assert_eq!(error.code, agent_client_protocol::ErrorCode::InvalidParams); + assert_eq!(error.code, -32602); let data = error .data - .as_ref() + .value() .and_then(|data| data.as_str()) .unwrap_or_default(); assert!(data.contains("protocolVersion"), "{error:?}"); @@ -967,21 +973,13 @@ fn sdk_supported_v2_method_surface_is_jsonrpc_mapped() -> Result<(), Error> { #[cfg(feature = "unstable_mcp_over_acp")] { - fn message_response() -> Result { - serde_json::from_value(serde_json::json!({ "tools": [] })) - .map_err(Error::into_internal_error) + fn message_response() -> v2::MessageMcpResponse { + v2::MessageMcpResponse::success(serde_json::json!({ "tools": [] })) } - assert_client_request!( - MessageMcpRequest, - MessageMcpResponse, - "mcp/message", - v2::MessageMcpRequest::new("connection-1", "tools/list"), - message_response()? - ); assert_v2_client_notification_mapping( "mcp/message", - v2::MessageMcpNotification::new("connection-1", "notifications/tools/list"), + v2::MessageMcpNotification::new("server-1", "request-1", "notifications/tools/list"), |notification| { matches!( notification, @@ -990,37 +988,13 @@ fn sdk_supported_v2_method_surface_is_jsonrpc_mapped() -> Result<(), Error> { }, )?; - assert_agent_request!( - ConnectMcpRequest, - ConnectMcpResponse, - "mcp/connect", - v2::ConnectMcpRequest::new("server-1"), - v2::ConnectMcpResponse::new("connection-1") - ); assert_agent_request!( MessageMcpRequest, MessageMcpResponse, "mcp/message", - v2::MessageMcpRequest::new("connection-1", "tools/list"), - message_response()? + v2::MessageMcpRequest::new("server-1", "request-1", "tools/list"), + message_response() ); - assert_agent_request!( - DisconnectMcpRequest, - DisconnectMcpResponse, - "mcp/disconnect", - v2::DisconnectMcpRequest::new("connection-1"), - v2::DisconnectMcpResponse::new() - ); - assert_v2_agent_notification_mapping( - "mcp/message", - v2::MessageMcpNotification::new("connection-1", "notifications/tools/list"), - |notification| { - matches!( - notification, - v2::AgentNotification::MessageMcpNotification(_) - ) - }, - )?; } let cancel_params = json_value(v2::CancelRequestNotification::new(String::from( @@ -1055,71 +1029,33 @@ fn mcp_over_acp_v1_variants_are_jsonrpc_mapped() -> Result<(), Error> { }}; } - assert_message_mapping!( - v1::ClientRequest, - "mcp/message", - json_value(v1::MessageMcpRequest::new("conn-1", "tools/list"))?, - v1::ClientRequest::MessageMcpRequest(_) - ); - assert_response_mapping!( - v1::AgentResponse, - "mcp/message", - serde_json::json!({ "tools": [] }), - v1::AgentResponse::MessageMcpResponse(_) - ); assert_message_mapping!( v1::ClientNotification, "mcp/message", json_value(v1::MessageMcpNotification::new( - "conn-1", + "server-1", + "request-1", "notifications/tools/list" ))?, v1::ClientNotification::MessageMcpNotification(_) ); - assert_message_mapping!( - v1::AgentRequest, - "mcp/connect", - json_value(v1::ConnectMcpRequest::new("server-1"))?, - v1::AgentRequest::ConnectMcpRequest(_) - ); assert_message_mapping!( v1::AgentRequest, "mcp/message", - json_value(v1::MessageMcpRequest::new("conn-1", "tools/list"))?, + json_value(v1::MessageMcpRequest::new( + "server-1", + "request-1", + "tools/list" + ))?, v1::AgentRequest::MessageMcpRequest(_) ); - assert_message_mapping!( - v1::AgentRequest, - "mcp/disconnect", - json_value(v1::DisconnectMcpRequest::new("conn-1"))?, - v1::AgentRequest::DisconnectMcpRequest(_) - ); - assert_response_mapping!( - v1::ClientResponse, - "mcp/connect", - json_value(v1::ConnectMcpResponse::new("conn-1"))?, - v1::ClientResponse::ConnectMcpResponse(_) - ); assert_response_mapping!( v1::ClientResponse, "mcp/message", - serde_json::json!({ "tools": [] }), - v1::ClientResponse::MessageMcpResponse(_) - ); - assert_response_mapping!( - v1::ClientResponse, - "mcp/disconnect", - serde_json::json!({}), - v1::ClientResponse::DisconnectMcpResponse(_) - ); - assert_message_mapping!( - v1::AgentNotification, - "mcp/message", - json_value(v1::MessageMcpNotification::new( - "conn-1", - "notifications/tools/list" + json_value(v1::MessageMcpResponse::success( + serde_json::json!({ "tools": [] }) ))?, - v1::AgentNotification::MessageMcpNotification(_) + v1::ClientResponse::MessageMcpResponse(_) ); Ok(()) @@ -2024,23 +1960,28 @@ async fn protocol_router_v2_only_rejects_v1_client() -> Result<(), Error> { channel .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "initialize".into(), json_value(v1_initialize_request(ProtocolVersion::V1))?, v1::RequestId::Number(1), )?)) + .await .map_err(Error::into_internal_error)?; while let Some(message) = channel.rx.next().await { - let TransportFrame::Single(message) = message else { + let TransportFrame::Single(message) = message.into_frame() else { continue; }; - let RawJsonRpcMessage::Response(v1::Response::Error { error, .. }) = message else { + let RawJsonRpcMessage::Response(agent_client_protocol::RawJsonRpcResponse::Error { + error, + .. + }) = message + else { continue; }; let data = error .data - .as_ref() + .value() .and_then(|data| data.as_str()) .unwrap_or_default(); assert!( @@ -2770,18 +2711,23 @@ async fn protocol_router_routes_future_protocol_version_to_v2() -> Result<(), Er ); channel .tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "initialize".into(), initialize, v1::RequestId::Number(1), )?)) + .await .map_err(Error::into_internal_error)?; while let Some(message) = channel.rx.next().await { - let TransportFrame::Single(message) = message else { + let TransportFrame::Single(message) = message.into_frame() else { continue; }; - let RawJsonRpcMessage::Response(v1::Response::Result { result, .. }) = message else { + let RawJsonRpcMessage::Response(agent_client_protocol::RawJsonRpcResponse::Result { + result, + .. + }) = message + else { continue; }; let initialize = v2::InitializeResponse::from_value("initialize", result)?; diff --git a/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs b/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs index fab29293..e5feea27 100644 --- a/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs +++ b/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs @@ -78,23 +78,31 @@ async fn request( let task = tokio::spawn(future); let request_id = v1::RequestId::Number(1); - tx.unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + tx.send_frame(TransportFrame::Single(RawJsonRpcMessage::request( method.into(), params, request_id.clone(), )?)) + .await .map_err(Error::into_internal_error)?; let result = loop { let frame = rx.next().await.ok_or_else(|| { Error::internal_error().data("proxy router closed before initialize response") })?; - let TransportFrame::Single(RawJsonRpcMessage::Response(response)) = frame else { + let TransportFrame::Single(RawJsonRpcMessage::Response(response)) = frame.into_frame() + else { continue; }; match response { - v1::Response::Result { id, result } if id == request_id => break Ok(result), - v1::Response::Error { id, error } if id == request_id => break Err(error), + agent_client_protocol::RawJsonRpcResponse::Result { id, result } + if id == request_id => + { + break Ok(result); + } + agent_client_protocol::RawJsonRpcResponse::Error { id, error } if id == request_id => { + break Err(error.into_acp_error()); + } _ => {} } }; diff --git a/src/agent-client-protocol/tests/session_ordering.rs b/src/agent-client-protocol/tests/session_ordering.rs index dee063af..4a88b9ec 100644 --- a/src/agent-client-protocol/tests/session_ordering.rs +++ b/src/agent-client-protocol/tests/session_ordering.rs @@ -1,8 +1,8 @@ use std::time::Duration; use agent_client_protocol::{ - ActiveSession, Agent, Channel, Client, Conductor, ConnectionTo, RawJsonRpcMessage, Responder, - SessionMessage, TransportBatch, TransportFrame, + ActiveSession, Agent, BudgetedFrame, Channel, Client, Conductor, ConnectionTo, + RawJsonRpcMessage, Responder, SessionMessage, TransportBatch, TransportFrame, schema::v1::{ ContentBlock, ContentChunk, NewSessionRequest, NewSessionResponse, PromptRequest, PromptResponse, SessionConfigOption, SessionConfigOptionCategory, @@ -34,16 +34,17 @@ async fn initialize_raw_v2_proxy( v2::Implementation::new(client_name, env!("CARGO_PKG_VERSION")), )); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "_proxy/initialize".to_owned(), serde_json::to_value(initialize).expect("initialize request should serialize"), initialize_id.clone(), )?)) + .await .expect("proxy should accept initialization"); let Some(TransportFrame::Single(RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { id, result }, - ))) = peer.rx.next().await + agent_client_protocol::RawJsonRpcResponse::Result { id, result }, + ))) = peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected the proxy initialize response"); }; @@ -306,7 +307,7 @@ async fn on_session_start_installs_routing_before_later_batch_entry() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a session/new request"); }; @@ -333,7 +334,8 @@ async fn on_session_start_installs_routing_before_later_batch_entry() { let batch = TransportBatch::from_messages([response, notification]) .expect("test batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept the response batch"); while peer.rx.next().await.is_some() {} @@ -407,16 +409,17 @@ async fn v2_proxy_session_start_installs_routing_before_later_batch_entry() { let upstream_id = agent_client_protocol::schema::v1::RequestId::Number(2); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "session/new".to_owned(), serde_json::to_value(v2::NewSessionRequest::new("/same-batch-v2-session")) .expect("session request should serialize"), upstream_id.clone(), )?)) + .await .expect("proxy should accept session/new"); let Some(TransportFrame::Single(RawJsonRpcMessage::Request(forwarded))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a forwarded session/new request"); }; @@ -448,18 +451,21 @@ async fn v2_proxy_session_start_installs_routing_before_later_batch_entry() { let batch = TransportBatch::from_messages([response, notification]) .expect("test response batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("proxy should accept the response batch"); let mut saw_response = false; let mut saw_update = false; for _ in 0..2 { - let Some(TransportFrame::Single(message)) = peer.rx.next().await else { + let Some(TransportFrame::Single(message)) = + peer.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected a forwarded session response and update"); }; match message { RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { id, result }, + agent_client_protocol::RawJsonRpcResponse::Result { id, result }, ) => { assert_eq!(id, upstream_id); let response = v2::NewSessionResponse::from_value("session/new", result)?; @@ -558,7 +564,7 @@ async fn v2_proxy_fork_installs_response_id_routing_before_later_batch_entry() { let upstream_id = agent_client_protocol::schema::v1::RequestId::Number(2); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "session/fork".to_owned(), serde_json::to_value(v2::ForkSessionRequest::new( source_session_id, @@ -567,10 +573,11 @@ async fn v2_proxy_fork_installs_response_id_routing_before_later_batch_entry() { .expect("fork request should serialize"), upstream_id.clone(), )?)) + .await .expect("proxy should accept session/fork"); let Some(TransportFrame::Single(RawJsonRpcMessage::Request(forwarded))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a forwarded session/fork request"); }; @@ -603,18 +610,21 @@ async fn v2_proxy_fork_installs_response_id_routing_before_later_batch_entry() { let batch = TransportBatch::from_messages([response, notification]) .expect("test response batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("proxy should accept the response batch"); let mut saw_response = false; let mut saw_update = false; for _ in 0..2 { - let Some(TransportFrame::Single(message)) = peer.rx.next().await else { + let Some(TransportFrame::Single(message)) = + peer.rx.next().await.map(BudgetedFrame::into_frame) + else { panic!("expected a forwarded fork response and update"); }; match message { RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { id, result }, + agent_client_protocol::RawJsonRpcResponse::Result { id, result }, ) => { assert_eq!(id, upstream_id); let response = v2::ForkSessionResponse::from_value("session/fork", result)?; @@ -714,7 +724,7 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { let upstream_id = agent_client_protocol::schema::v1::RequestId::Number(2); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "session/resume".to_owned(), serde_json::to_value(v2::ResumeSessionRequest::new( session_id.clone(), @@ -723,10 +733,11 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { .expect("resume request should serialize"), upstream_id.clone(), )?)) + .await .expect("proxy should accept session/resume"); let Some(TransportFrame::Single(RawJsonRpcMessage::Request(forwarded))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected a forwarded session/resume request"); }; @@ -755,11 +766,12 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { let batch = TransportBatch::from_messages([notification, response]) .expect("test replay batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("proxy should accept the replay batch"); let Some(TransportFrame::Single(RawJsonRpcMessage::Notification(notification))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected replay to be forwarded before the resume response"); }; @@ -774,8 +786,8 @@ async fn v2_proxy_resume_forwards_replay_before_same_batch_response() { )); let Some(TransportFrame::Single(RawJsonRpcMessage::Response( - agent_client_protocol::schema::v1::Response::Result { id, result }, - ))) = peer.rx.next().await + agent_client_protocol::RawJsonRpcResponse::Result { id, result }, + ))) = peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected the resume response after replay"); }; diff --git a/src/agent-client-protocol/tests/session_restore.rs b/src/agent-client-protocol/tests/session_restore.rs index c01875e0..6938d5b5 100644 --- a/src/agent-client-protocol/tests/session_restore.rs +++ b/src/agent-client-protocol/tests/session_restore.rs @@ -4,8 +4,9 @@ use std::{future::pending, path::PathBuf, time::Duration}; use agent_client_protocol::{ - Agent, Channel, Client, ConnectionTo, Error, ErrorCode, JsonRpcMessage, JsonRpcNotification, - RawJsonRpcMessage, Responder, SessionMessage, TransportBatch, TransportFrame, UntypedMessage, + Agent, BudgetedFrame, Channel, Client, ConnectionTo, Error, ErrorCode, JsonRpcMessage, + JsonRpcNotification, RawJsonRpcMessage, Responder, SessionMessage, TransportBatch, + TransportFrame, UntypedMessage, schema::v1::{ CancelRequestNotification, ContentBlock, ContentChunk, LoadSessionRequest, LoadSessionResponse, RequestId, ResumeSessionRequest, ResumeSessionResponse, @@ -138,7 +139,7 @@ async fn load_session_preserves_pre_response_replay_and_exact_response() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/load") }; @@ -153,7 +154,8 @@ async fn load_session_preserves_pre_response_replay_and_exact_response() { ]) .expect("restore batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept replay and response"); while peer.rx.next().await.is_some() {} @@ -206,7 +208,7 @@ async fn resume_session_returns_exact_response_and_an_active_session() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/resume") }; @@ -219,7 +221,8 @@ async fn resume_session_returns_exact_response_and_an_active_session() { ]) .expect("resume batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept update and response"); while peer.rx.next().await.is_some() {} @@ -256,7 +259,7 @@ async fn resume_session_from_preserves_the_existing_request() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/resume") }; @@ -265,10 +268,11 @@ async fn resume_session_from_preserves_the_existing_request() { peer_request ); peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Ok(serde_json::to_value(ResumeSessionResponse::new())?), ))) + .await .expect("client should accept resume response"); while peer.rx.next().await.is_some() {} Ok::<(), Error>(()) @@ -314,11 +318,12 @@ async fn restore_waits_for_routing_acknowledgment_before_publication() { let peer = async move { peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::request( "test/block-incoming".to_owned(), serde_json::json!({}), RequestId::Number(1), )?)) + .await .expect("client should accept the blocking request"); restore_called_rx .await @@ -381,7 +386,7 @@ async fn failed_restore_removes_routing_before_later_batch_entries() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/load") }; @@ -395,18 +400,21 @@ async fn failed_restore_removes_routing_before_later_batch_entries() { ]) .expect("failure batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(batch)) + .send_frame(TransportFrame::Batch(batch)) + .await .expect("client should accept failure, probe, and barrier"); - let Some(TransportFrame::Single(RawJsonRpcMessage::Request(retry))) = peer.rx.next().await + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(retry))) = + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected retry session/load") }; peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( retry.id, Ok(serde_json::to_value(LoadSessionResponse::new())?), ))) + .await .expect("client should accept retry response"); while peer.rx.next().await.is_some() {} Ok::<(), Error>(()) @@ -477,7 +485,7 @@ async fn cancelling_restore_cancels_request_and_removes_routing() { let peer = async move { let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected session/resume") }; @@ -490,7 +498,7 @@ async fn cancelling_restore_cancels_request_and_removes_routing() { .map_err(Error::into_internal_error)?; let Some(TransportFrame::Single(RawJsonRpcMessage::Notification(notification))) = - peer.rx.next().await + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("dropping the restore future should send $/cancel_request") }; @@ -504,28 +512,32 @@ async fn cancelling_restore_cancels_request_and_removes_routing() { ]) .expect("cancellation probe batch should be non-empty"); peer.tx - .unbounded_send(TransportFrame::Batch(probe_batch)) + .send_frame(TransportFrame::Batch(probe_batch)) + .await .expect("client should accept cancellation probe and barrier"); barrier_observed_rx .await .map_err(Error::into_internal_error)?; peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( request.id, Err(Error::request_cancelled()), ))) + .await .expect("client should accept the cancelled request's response"); - let Some(TransportFrame::Single(RawJsonRpcMessage::Request(retry))) = peer.rx.next().await + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(retry))) = + peer.rx.next().await.map(BudgetedFrame::into_frame) else { panic!("expected retry session/resume") }; peer.tx - .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + .send_frame(TransportFrame::Single(RawJsonRpcMessage::response( retry.id, Ok(serde_json::to_value(ResumeSessionResponse::new())?), ))) + .await .expect("client should accept retry response"); while peer.rx.next().await.is_some() {} Ok::<(), Error>(()) diff --git a/src/agent-client-protocol/tests/session_v2_mcp.rs b/src/agent-client-protocol/tests/session_v2_mcp.rs index 904e7f52..81009032 100644 --- a/src/agent-client-protocol/tests/session_v2_mcp.rs +++ b/src/agent-client-protocol/tests/session_v2_mcp.rs @@ -12,8 +12,8 @@ use std::{ }; use agent_client_protocol::{ - Agent, Client, ConnectTo, ConnectionTo, DynConnectTo, Error, ErrorCode, JsonRpcNotification, - JsonRpcRequest, JsonRpcResponse, Responder, RunWithConnectionTo, V2ConnectionTo, + Agent, Client, ConnectTo, ConnectionTo, DynConnectTo, Error, ErrorCode, JsonRpcRequest, + JsonRpcResponse, Responder, RunWithConnectionTo, V2ConnectionTo, mcp_server::{McpConnectionTo, McpServer, McpServerConnect}, role, schema::{ProtocolVersion, v2}, @@ -87,21 +87,14 @@ struct ConnectionProbeResponse { nonce: String, } -#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcNotification)] -#[notification(method = "_test/notice")] -struct NoticeNotification { - message: String, -} - #[derive(Debug, PartialEq, Eq)] struct ObservedMcpContext { server_id: String, - connection_id: String, + request_id: String, } struct EchoMcpConnect { context_tx: mpsc::UnboundedSender, - notice_tx: mpsc::UnboundedSender, runner_started: Arc, dropped_tx: Mutex>>, } @@ -135,37 +128,23 @@ impl McpServerConnect for EchoMcpConnect { .server_id() .expect("the MCP server should be attached through ACP") .to_string(), - connection_id: context - .connection_id() - .expect("an attached MCP connection should have an ID") + request_id: context + .request_id() + .expect("an attached MCP request should have an ID") .to_string(), }) .expect("MCP context receiver should remain active"); - DynConnectTo::new(EchoMcpComponent { - notice_tx: self.notice_tx.clone(), - }) + DynConnectTo::new(EchoMcpComponent) } } -struct EchoMcpComponent { - notice_tx: mpsc::UnboundedSender, -} +struct EchoMcpComponent; impl ConnectTo for EchoMcpComponent { async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { - let notice_tx = self.notice_tx; - role::mcp::Server .builder() - .on_receive_notification( - async move |notification: NoticeNotification, _connection| { - notice_tx - .unbounded_send(notification.message) - .map_err(Error::into_internal_error) - }, - agent_client_protocol::on_receive_notification!(), - ) .on_receive_request( async |request: EchoRequest, responder: Responder, _connection| { responder.respond(EchoResponse { @@ -210,8 +189,7 @@ impl RunWithConnectionTo for ProbeRunner { #[derive(Debug)] struct RoundTrip { server_id: String, - connection_id: String, - notice: String, + request_id: String, response: Value, } @@ -220,36 +198,30 @@ async fn run_mcp_round_trip( server_id: &v2::McpServerAcpId, sequence: usize, ) -> Result { - let connected = connection - .send_request(v2::ConnectMcpRequest::new(server_id.clone())) - .block_task() - .await?; - let connection_id = connected.connection_id; - let notice = format!("notice-{sequence}"); - connection.send_notification( - v2::MessageMcpNotification::new(connection_id.clone(), "_test/notice") - .params(object(json!({ "message": notice }))), - )?; - + let request_id = format!("request-{sequence}"); let message = format!("message-{sequence}"); let response = connection .send_request( - v2::MessageMcpRequest::new(connection_id.clone(), "_test/echo") - .params(object(json!({ "message": message }))), + v2::MessageMcpRequest::new(server_id.clone(), request_id.clone(), "_test/echo").params( + object(json!({ "message": message, "_meta": { + "io.modelcontextprotocol/protocolVersion": "2026-07-28", + "io.modelcontextprotocol/clientCapabilities": {} + } })), + ), ) .block_task() .await?; - let response = serde_json::from_str(response.0.get()).map_err(Error::into_internal_error)?; - - connection - .send_request(v2::DisconnectMcpRequest::new(connection_id.clone())) - .block_task() - .await?; + let response = match response { + v2::MessageMcpResponse::Result { result, .. } => result, + v2::MessageMcpResponse::Error { error, .. } => { + return Err(Error::new(error.code, error.message)); + } + _ => return Err(Error::internal_error().data("unknown MCP response carrier")), + }; Ok(RoundTrip { server_id: server_id.to_string(), - connection_id: connection_id.to_string(), - notice, + request_id, response, }) } @@ -258,15 +230,12 @@ async fn assert_round_trip( sequence: usize, round_trip_rx: &mut UnboundedReceiver>, context_rx: &mut UnboundedReceiver, - notice_rx: &mut UnboundedReceiver, ) -> Result<(), Error> { let round_trip = next(round_trip_rx, "MCP round trip").await?; - let context = next(context_rx, "MCP connection context").await; - let notice = next(notice_rx, "inner MCP notification").await; + let context = next(context_rx, "MCP request context").await; assert_eq!(context.server_id, round_trip.server_id); - assert_eq!(context.connection_id, round_trip.connection_id); - assert_eq!(notice, round_trip.notice); + assert_eq!(context.request_id, round_trip.request_id); assert_eq!( round_trip.response, json!({ "echoed": format!("message-{sequence}") }) @@ -364,7 +333,6 @@ async fn v2_session_mcp_attachment_is_ready_during_setup_and_lives_for_connectio let test = async move { let (context_tx, mut context_rx) = mpsc::unbounded(); - let (notice_tx, mut notice_rx) = mpsc::unbounded(); let (connector_dropped_tx, connector_dropped_rx) = oneshot::channel(); let (runner_started_tx, runner_started_rx) = oneshot::channel(); let (runner_dropped_tx, runner_dropped_rx) = oneshot::channel(); @@ -384,7 +352,6 @@ async fn v2_session_mcp_attachment_is_ready_during_setup_and_lives_for_connectio let mcp_server = McpServer::::new( EchoMcpConnect { context_tx, - notice_tx, runner_started: runner_started.clone(), dropped_tx: Mutex::new(Some(connector_dropped_tx)), }, @@ -411,7 +378,7 @@ async fn v2_session_mcp_attachment_is_ready_during_setup_and_lives_for_connectio "the MCP runner must be first-polled before session/new is published" ); - assert_round_trip(1, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(1, &mut round_trip_rx, &mut context_rx).await?; let session = pending_session.block_task().await?.into_session(); let remaining_session = session.clone(); @@ -421,7 +388,7 @@ async fn v2_session_mcp_attachment_is_ready_during_setup_and_lives_for_connectio round_trip_trigger_tx .unbounded_send(()) .map_err(Error::into_internal_error)?; - assert_round_trip(2, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(2, &mut round_trip_rx, &mut context_rx).await?; Ok(()) }) @@ -550,7 +517,6 @@ async fn v2_fork_mcp_attachment_preserves_request_and_lives_for_connection() -> let test = async move { let (context_tx, mut context_rx) = mpsc::unbounded(); - let (notice_tx, mut notice_rx) = mpsc::unbounded(); let (connector_dropped_tx, connector_dropped_rx) = oneshot::channel(); let (runner_started_tx, runner_started_rx) = oneshot::channel(); let (runner_dropped_tx, runner_dropped_rx) = oneshot::channel(); @@ -570,7 +536,6 @@ async fn v2_fork_mcp_attachment_preserves_request_and_lives_for_connection() -> let mcp_server = McpServer::::new( EchoMcpConnect { context_tx, - notice_tx, runner_started: runner_started.clone(), dropped_tx: Mutex::new(Some(connector_dropped_tx)), }, @@ -598,7 +563,7 @@ async fn v2_fork_mcp_attachment_preserves_request_and_lives_for_connection() -> "the MCP runner must be first-polled before session/fork is published" ); - assert_round_trip(1, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(1, &mut round_trip_rx, &mut context_rx).await?; let opened = pending_session.block_task().await?; assert_eq!(opened.session().session_id(), &expected_forked_session_id); @@ -612,7 +577,7 @@ async fn v2_fork_mcp_attachment_preserves_request_and_lives_for_connection() -> round_trip_trigger_tx .unbounded_send(()) .map_err(Error::into_internal_error)?; - assert_round_trip(2, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(2, &mut round_trip_rx, &mut context_rx).await?; Ok(()) }) @@ -742,7 +707,6 @@ async fn v2_resume_mcp_attachment_preserves_request_and_lives_for_connection() - let test = async move { let (context_tx, mut context_rx) = mpsc::unbounded(); - let (notice_tx, mut notice_rx) = mpsc::unbounded(); let (connector_dropped_tx, connector_dropped_rx) = oneshot::channel(); let (runner_started_tx, runner_started_rx) = oneshot::channel(); let (runner_dropped_tx, runner_dropped_rx) = oneshot::channel(); @@ -762,7 +726,6 @@ async fn v2_resume_mcp_attachment_preserves_request_and_lives_for_connection() - let mcp_server = McpServer::::new( EchoMcpConnect { context_tx, - notice_tx, runner_started: runner_started.clone(), dropped_tx: Mutex::new(Some(connector_dropped_tx)), }, @@ -793,7 +756,7 @@ async fn v2_resume_mcp_attachment_preserves_request_and_lives_for_connection() - "the MCP runner must be first-polled before session/resume is published" ); - assert_round_trip(1, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(1, &mut round_trip_rx, &mut context_rx).await?; let opened = pending_session.block_task().await?; assert_eq!(opened.session().session_id(), &session_id); @@ -806,7 +769,7 @@ async fn v2_resume_mcp_attachment_preserves_request_and_lives_for_connection() - round_trip_trigger_tx .unbounded_send(()) .map_err(Error::into_internal_error)?; - assert_round_trip(2, &mut round_trip_rx, &mut context_rx, &mut notice_rx).await?; + assert_round_trip(2, &mut round_trip_rx, &mut context_rx).await?; Ok(()) })