From 1025c64a6aeebe162efa51f4247d2f11d8487b01 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Thu, 1 Oct 2026 18:38:08 +0200 Subject: [PATCH 1/3] fix(acp)!: model connection-driver lifetimes explicitly Preserve passive half-closes and poll owned work alongside frame forwarding. Drain accepted output before completion using private built-in transport finish coordination, and update protocol and HTTP consumers to preserve driver identity. Reject escaped output producers on active completion and drain the installed HTTP router before unregistering. Add lifecycle regressions, migration guidance, and focused changelog entries without resource-policy changes. BREAKING CHANGE: ConnectTo::into_channel_and_future now returns (Channel, ConnectionDriver) instead of a channel and boxed future. Custom overrides must mark owned versus passive work; existing raw Channel APIs remain unchanged. --- md/SUMMARY.md | 1 + md/migration-connection-drivers.md | 72 ++++ md/transport-architecture.md | 40 ++- src/agent-client-protocol-http/CHANGELOG.md | 11 + src/agent-client-protocol-http/src/client.rs | 6 +- .../src/connection.rs | 221 ++++++++---- .../src/http_server.rs | 51 +-- .../src/websocket_server.rs | 41 +-- .../src/mcp_over_acp/http.rs | 14 +- src/agent-client-protocol/CHANGELOG.md | 15 + src/agent-client-protocol/src/acp_agent.rs | 2 +- src/agent-client-protocol/src/component.rs | 135 ++++++- src/agent-client-protocol/src/jsonrpc.rs | 238 ++++++++---- .../src/jsonrpc/transport_actor.rs | 2 +- src/agent-client-protocol/src/lib.rs | 2 +- src/agent-client-protocol/src/role/acp.rs | 340 ++++++++++++++++-- .../tests/jsonrpc_transport_close.rs | 307 +++++++++++++++- 17 files changed, 1229 insertions(+), 269 deletions(-) create mode 100644 md/migration-connection-drivers.md diff --git a/md/SUMMARY.md b/md/SUMMARY.md index c51ce8a0..ba11a00e 100644 --- a/md/SUMMARY.md +++ b/md/SUMMARY.md @@ -31,6 +31,7 @@ # Reference +- [Migrating Connection Drivers](./migration-connection-drivers.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/migration-connection-drivers.md b/md/migration-connection-drivers.md new file mode 100644 index 00000000..d55d7a25 --- /dev/null +++ b/md/migration-connection-drivers.md @@ -0,0 +1,72 @@ +# Migrating Connection Drivers + +`ConnectTo::into_channel_and_future` now returns `(Channel, ConnectionDriver)` +instead of `(Channel, BoxFuture<'static, Result<()>>)`. The same change applies +when accessing a component through `DynConnectTo`. + +This is a source-breaking transport-adapter change. It does not change ACP wire +messages, the raw `Channel` sender/receiver types, or `unbounded_send`, and it +introduces no new frame-size, queue, or task limits. + +## Components using the default conversion + +If your component implements only `connect_to`, no change is needed. The +default conversion still creates a channel pair and drives your component, now +returning an owned `ConnectionDriver`. + +Existing callers that infer the returned type and await the driver continue +to work: + +```rust,ignore +let (channel, driver) = component.into_channel_and_future(); +// Use channel while continuing to poll the driver. +driver.await?; +``` + +## Custom conversion overrides + +Import `ConnectionDriver` from `agent_client_protocol` and change the return +type. Wrap a future that owns the connection work with `ConnectionDriver::new`: + +```rust,ignore +fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { + let (channel, future) = self.into_channel_transport(); + (channel, ConnectionDriver::new(future)) +} +``` + +For an endpoint whose work is driven elsewhere, return +`ConnectionDriver::passive()` instead of wrapping a ready no-op future: + +```rust,ignore +fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { + (self.channel, ConnectionDriver::passive()) +} +``` + +Passive drivers are awaitable and immediately succeed, but that success is +**not EOF**. Bridges must check `is_passive()` before using driver completion +as a lifetime signal. Passive bridges retain each read/write half until its +own closure; an input half-close can still be followed by a final response. + +If a wrapper simply exposes another component's endpoint, return its original +`(channel, driver)` pair. Re-boxing that driver and wrapping it with `new` +would erase the passive distinction and any built-in finish coordination. + +## Completion and drain responsibilities + +Poll owned work and outbound forwarding concurrently. A driver may need its +outbound request to be delivered before it can receive a response and finish. + +An adapter must not report success before flushing output it already accepted. +Custom normalized drivers should finish their own work and drain output after +their channel input closes. The SDK cannot infer how to flush an arbitrary +opaque future or external buffer. + +Built-in `Lines` and `ByteStreams` preserve normal half-close behavior. When +their owner explicitly finishes, they drain accepted output while continuing +to poll incoming I/O for errors, rather than waiting for unrelated remote +input to reach EOF. Errors may terminate the connection without graceful drain. + +See [Transport Architecture](./transport-architecture.md#component-boundary) +for the active/passive boundary and forwarding rules. diff --git a/md/transport-architecture.md b/md/transport-architecture.md index 8204b77e..08d72a19 100644 --- a/md/transport-architecture.md +++ b/md/transport-architecture.md @@ -45,7 +45,7 @@ by the JSON-RPC envelope types from `agent-client-protocol-schema`: enum RawJsonRpcMessage { Request(Request), Notification(Notification), - Response(Response), + Response(RawJsonRpcResponse), } ``` @@ -262,16 +262,42 @@ 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. +The channel carries only `TransportFrame` wire events. The awaitable +`ConnectionDriver` distinguishes owned work from a passive endpoint: + +| Driver | Meaning of successful completion | +| --- | --- | +| `ConnectionDriver::new(future)` | The component's owned work has finished. | +| `ConnectionDriver::passive()` | No work is owned here; readiness says nothing about either I/O half. | + +A raw `Channel` is passive. Its bridge preserves both directions independently: +one sender closing must not prevent a final response in the reverse direction. +Owned completion lets a bridge stop accepting new output, drain frames already +accepted, and finish without waiting for unrelated remote input to close. +Outbound forwarding must remain polled while owned work is running; otherwise +a component waiting for a response to its own request could deadlock. + +Buffered adapters are responsible for flushing their accepted output before +reporting successful completion. The built-in line and byte-stream adapters +keep the read half moving during write drain and propagate incoming errors; +their explicit finish handling does not require remote read EOF. Merely +wrapping an arbitrary future cannot make an opaque custom adapter drain safely. + +Most components implement only `connect_to`; default normalization supplies +the owned driver. Direct transports override `into_channel_and_future` to avoid +an intermediate copy. Wrappers that expose an existing endpoint should forward +its driver unchanged so passive identity and built-in completion handling are +not lost. + +See [Migrating Connection Drivers](./migration-connection-drivers.md) for custom +override changes. This lifecycle distinction does not change the existing raw +channel types or introduce frame-size, queue, or task limits. ## Transport Implementations diff --git a/src/agent-client-protocol-http/CHANGELOG.md b/src/agent-client-protocol-http/CHANGELOG.md index 40dc362a..44c21d23 100644 --- a/src/agent-client-protocol-http/CHANGELOG.md +++ b/src/agent-client-protocol-http/CHANGELOG.md @@ -2,11 +2,22 @@ ## [Unreleased] +### Changed + +- Adapt `HttpClient`'s `ConnectTo` conversion to the core SDK's breaking + `ConnectionDriver` return type. Channels and HTTP framing remain unchanged; + no new resource limits are introduced. + ### Fixed - Preserve raw JSON-RPC error codes, omitted versus null data, and error extension fields across HTTP/SSE and WebSocket transports, using the core SDK's new `RawJsonRpcResponse` representation. +- Do not treat a passive agent-factory endpoint's no-op driver as agent + completion while its transport remains open. +- On active agent completion, reject further output from escaped sender clones + and drain accepted frames before removing the connection and closing its + streams, without waiting for those clones to be dropped. ## [2.2.0](https://github.com/agentclientprotocol/rust-sdk/compare/agent-client-protocol-http-v2.1.0...agent-client-protocol-http-v2.2.0) - 2026-09-18 diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index 86085762..6af4cd73 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -4,7 +4,7 @@ use std::{ }; use agent_client_protocol::{ - Agent, Channel, Client, ConnectTo, Error as AcpError, RawJsonRpcMessage, + Agent, Channel, Client, ConnectTo, ConnectionDriver, Error as AcpError, RawJsonRpcMessage, RawJsonRpcResponse as RpcResponse, TransportBatchEntry, TransportFrame, schema::v1::RequestId, }; use async_tungstenite::tungstenite::Message as WsMessage; @@ -122,9 +122,9 @@ impl ConnectTo for HttpClient { } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), AcpError>>) { + fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { let (caller, transport) = Channel::duplex(); - (caller, Box::pin(run(self, transport))) + (caller, ConnectionDriver::new(run(self, transport))) } } diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index f2796bab..cd635b92 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -1,13 +1,14 @@ use std::{ collections::{HashMap, VecDeque}, sync::{Arc, Mutex as StdMutex, Weak}, + task::Poll, }; use agent_client_protocol::{ Channel, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame, schema::v1::RequestId, }; -use futures::{SinkExt, StreamExt}; +use futures::{FutureExt, SinkExt, StreamExt}; use tokio::sync::{Mutex, RwLock, mpsc, watch}; use tracing::{debug, error, trace}; @@ -189,7 +190,7 @@ impl Connection { pub(crate) async fn shutdown(&self) { // Explicit peer teardown is abortive. Natural agent completion instead - // awaits the router in `close_connection_task` before closing streams. + // awaits its router before unregistering and closing streams. self.close_streams(); if let Some(h) = self.agent_handle.lock().await.take() { h.abort(); @@ -423,12 +424,7 @@ pub(crate) struct ConnectionRegistry { } pub(crate) trait AgentFactory: Send + Sync + 'static { - fn spawn_agent( - &self, - ) -> ( - Channel, - futures::future::BoxFuture<'static, agent_client_protocol::Result<()>>, - ); + fn spawn_agent(&self) -> (Channel, agent_client_protocol::ConnectionDriver); } impl AgentFactory for F @@ -436,12 +432,7 @@ where F: Fn() -> C + Send + Sync + 'static, C: agent_client_protocol::ConnectTo, { - fn spawn_agent( - &self, - ) -> ( - Channel, - futures::future::BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, agent_client_protocol::ConnectionDriver) { self().into_channel_and_future() } } @@ -502,8 +493,24 @@ 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 (finish_outbound_tx, finish_outbound_rx) = futures::channel::oneshot::channel::<()>(); let outbound = async move { - while let Some(msg) = agent_rx.next().await { + let mut finish_outbound_rx = Some(finish_outbound_rx); + while let Some(msg) = futures::future::poll_fn(|cx| { + if let Some(finish) = &mut finish_outbound_rx + && let Poll::Ready(result) = finish.poll_unpin(cx) + { + finish_outbound_rx = None; + if result.is_ok() { + // Active completion rejects escaped producers while + // leaving already accepted frames available to drain. + agent_rx.close(); + } + } + agent_rx.poll_next_unpin(cx) + }) + .await + { if outbound_tx.send(msg).is_err() { inbound_abort_for_outbound.abort(); break; @@ -534,6 +541,11 @@ impl ConnectionRegistry { let agent_handle = tokio::spawn(async move { let conn_id_for_agent = conn_id_for_task.clone(); let agent = async move { + if agent_future.is_passive() { + // A passive endpoint is driven only by the channel pumps; + // its immediately ready no-op driver is not connection EOF. + std::future::pending::<()>().await; + } if let Err(e) = agent_future.await { error!(connection_id = %conn_id_for_agent, "ACP agent task error: {e}"); } @@ -543,13 +555,17 @@ impl ConnectionRegistry { match futures::future::select(agent, pump).await { futures::future::Either::Left(((), pump)) => { inbound_abort.abort(); + let _sent = finish_outbound_tx.send(()); pump.await; } futures::future::Either::Right(((), _agent)) => {} } debug!(connection_id = %conn_id_for_task, "ACP connection task ended"); + let connection_to_close = drain_connection_router(connection_for_task).await; connections.write().await.remove(&conn_id_for_task); - close_connection_task(connection_for_task).await; + if let Some(connection) = connection_to_close { + connection.close_streams(); + } }); *connection.agent_handle.lock().await = Some(agent_handle); @@ -571,17 +587,15 @@ impl ConnectionRegistry { } } -async fn close_connection_task(connection: Weak) { - let Some(connection) = connection.upgrade() else { - return; - }; +async fn drain_connection_router(connection: Weak) -> Option> { + let connection = connection.upgrade()?; let router_handle = connection.router_handle.lock().await.take(); if let Some(h) = router_handle && let Err(error) = h.await { error!("outbound router task failed while draining: {error}"); } - connection.close_streams(); + Some(connection) } fn pending_route_key(id: &RequestId) -> Option { @@ -608,8 +622,7 @@ fn take_pending_route( mod tests { use std::sync::Arc; - use agent_client_protocol::TransportBatch; - use futures::future::BoxFuture; + use agent_client_protocol::{ConnectionDriver, TransportBatch}; use tokio::{ sync::Notify, time::{Duration, sleep, timeout}, @@ -741,15 +754,10 @@ mod tests { } impl AgentFactory for ExitingAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); let exit = self.exit.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { exit.notified().await; drop(agent); Ok(()) @@ -762,14 +770,9 @@ mod tests { struct RespondThenExitAgentFactory; impl AgentFactory for RespondThenExitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { agent .tx .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( @@ -789,15 +792,10 @@ mod tests { } impl AgentFactory for MalformedThenWaitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); let emit = self.emit.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { emit.notified().await; agent .tx @@ -820,16 +818,11 @@ mod tests { } impl AgentFactory for SendThenWaitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); let message = self.message.clone(); let exit = self.exit.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { agent .tx .unbounded_send(TransportFrame::Single(message)) @@ -847,15 +840,10 @@ mod tests { } impl AgentFactory for BatchThenWaitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); let exit = self.exit.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { let batch = TransportBatch::from_messages([ RawJsonRpcMessage::notification( "test/first".to_string(), @@ -883,18 +871,16 @@ mod tests { struct FinalFrameThenExitAgentFactory { emit: Arc, + escaped_output: + Arc>>>, } impl AgentFactory for FinalFrameThenExitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); + *self.escaped_output.lock().unwrap() = Some(agent.tx.clone()); let emit = self.emit.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { emit.notified().await; agent .tx @@ -913,6 +899,36 @@ mod tests { } } + #[tokio::test] + async fn passive_agent_driver_does_not_close_http_channel_pumps() { + let (endpoint, mut remote) = Channel::duplex(); + let endpoint = std::sync::Mutex::new(Some(endpoint)); + let registry = ConnectionRegistry::new(Arc::new(move || { + endpoint.lock().unwrap().take().expect("one connection") + })); + let (connection_id, connection) = registry.create_connection().await; + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/passive".into(), serde_json::json!({})).unwrap(), + ); + connection.inbound_tx.send(frame.clone()).unwrap(); + assert!( + timeout(Duration::from_secs(1), remote.rx.next()) + .await + .expect("passive readiness must not abort inbound forwarding") + .is_some() + ); + remote.tx.unbounded_send(frame).unwrap(); + assert!( + timeout(Duration::from_secs(1), connection.recv_initial()) + .await + .expect("passive endpoint must keep its reverse direction alive") + .is_some() + ); + assert!(registry.get(&connection_id).await.is_some()); + connection.shutdown().await; + registry.remove(&connection_id).await; + } + #[tokio::test] async fn agent_exit_removes_connection_and_closes_streams() { let exit = Arc::new(Notify::new()); @@ -998,12 +1014,15 @@ mod tests { #[tokio::test] async fn agent_exit_flushes_final_frame_before_closing_streams() { let emit = Arc::new(Notify::new()); + let escaped_output = Arc::new(StdMutex::new(None)); let registry = ConnectionRegistry::new(Arc::new(FinalFrameThenExitAgentFactory { emit: emit.clone(), + escaped_output: escaped_output.clone(), })); let (connection_id, connection) = registry.create_connection().await; let mut outbound = connection.subscribe_connection_stream().unwrap(); connection.start_router().await; + let escaped_output = escaped_output.lock().unwrap().take().unwrap(); emit.notify_one(); timeout(Duration::from_secs(1), async { @@ -1026,6 +1045,78 @@ mod tests { RawJsonRpcMessage::Notification(notification) if notification.method.as_ref() == "test/final" )); + assert!( + escaped_output + .unbounded_send(TransportFrame::Single( + RawJsonRpcMessage::notification("test/late".into(), serde_json::json!({})) + .unwrap(), + )) + .is_err(), + "active completion must reject escaped senders without dropping them" + ); + // Stream termination is signalled by closed_tx; the connection keeps + // its mailbox sender alive so later subscribers can drain queued data. + assert_eq!(outbound.try_recv(), Err(mpsc::error::TryRecvError::Empty)); + } + + #[tokio::test] + async fn active_agent_remains_registered_until_outbound_router_drains() { + let emit = Arc::new(Notify::new()); + let registry = ConnectionRegistry::new(Arc::new(FinalFrameThenExitAgentFactory { + emit: emit.clone(), + escaped_output: Arc::new(StdMutex::new(None)), + })); + let (connection_id, connection) = registry.create_connection().await; + let mut outbound = connection.subscribe_connection_stream().unwrap(); + let mut frames = connection.outbound_rx.lock().await.take().unwrap(); + let (routing_started_tx, routing_started_rx) = tokio::sync::oneshot::channel(); + let (release_router_tx, release_router_rx) = tokio::sync::oneshot::channel(); + let routing_connection = connection.clone(); + *connection.router_handle.lock().await = Some(tokio::spawn(async move { + let frame = frames.recv().await.expect("accepted final frame"); + let _sent = routing_started_tx.send(()); + // A dropped release sender also unblocks failed-test cleanup. + let _released = release_router_rx.await; + routing_connection.route_outbound(frame).await.unwrap(); + while let Some(frame) = frames.recv().await { + routing_connection.route_outbound(frame).await.unwrap(); + } + })); + + emit.notify_one(); + timeout(Duration::from_secs(1), async { + routing_started_rx.await.unwrap(); + // Taking the join handle establishes that natural shutdown has + // reached router drain, not merely that the router is scheduled. + while connection.router_handle.lock().await.is_some() { + tokio::task::yield_now().await; + } + }) + .await + .expect("natural shutdown should await the gated router"); + assert!( + registry.get(&connection_id).await.is_some(), + "the connection must remain discoverable until accepted output is routed" + ); + let mut closed = connection.subscribe_closed(); + assert!(!*closed.borrow()); + + release_router_tx.send(()).unwrap(); + timeout(Duration::from_secs(1), async { + while !*closed.borrow() { + closed.changed().await.unwrap(); + } + }) + .await + .expect("closure should follow router drain and registry removal"); + assert!(registry.get(&connection_id).await.is_none()); + let text = outbound.try_recv().expect("the final frame must be routed"); + let message = serde_json::from_str::(&text).unwrap(); + assert!(matches!( + message, + RawJsonRpcMessage::Notification(notification) + if notification.method.as_ref() == "test/final" + )); } #[tokio::test] diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index 4f3ef608..5cbc7479 100644 --- a/src/agent-client-protocol-http/src/http_server.rs +++ b/src/agent-client-protocol-http/src/http_server.rs @@ -480,10 +480,10 @@ mod tests { use std::sync::Arc; use agent_client_protocol::{ - Channel, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, TransportFrame, - schema::v1::RequestId, + Channel, ConnectionDriver, RawJsonRpcMessage, TransportBatch, TransportBatchEntry, + TransportFrame, schema::v1::RequestId, }; - use futures::{StreamExt, future::BoxFuture}; + use futures::StreamExt; use serde_json::json; use tokio::{ sync::mpsc, @@ -500,15 +500,10 @@ mod tests { } impl AgentFactory for CapturingAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { let Channel { rx: mut incoming, tx: _, @@ -531,14 +526,9 @@ mod tests { struct RejectingInitializeAgentFactory; impl AgentFactory for RejectingInitializeAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (mut agent, transport) = Channel::duplex(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { match agent.rx.next().await { Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) => { agent @@ -585,14 +575,9 @@ mod tests { struct PendingInitializeAgentFactory; impl AgentFactory for PendingInitializeAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { let Channel { rx: mut incoming, tx: _outgoing, @@ -610,15 +595,10 @@ mod tests { } impl AgentFactory for BatchAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (mut agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { panic!("expected one batch frame"); }; @@ -678,14 +658,9 @@ mod tests { struct SideTrafficBeforeInitializeResponseAgentFactory; impl AgentFactory for SideTrafficBeforeInitializeResponseAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (mut agent, transport) = Channel::duplex(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { panic!("expected one initial batch frame"); }; diff --git a/src/agent-client-protocol-http/src/websocket_server.rs b/src/agent-client-protocol-http/src/websocket_server.rs index faf57c60..1ab3b8ab 100644 --- a/src/agent-client-protocol-http/src/websocket_server.rs +++ b/src/agent-client-protocol-http/src/websocket_server.rs @@ -235,12 +235,11 @@ where #[cfg(test)] mod tests { use agent_client_protocol::{ - Channel, RawJsonRpcResponse as RpcResponse, TransportBatch, TransportBatchEntry, - TransportFrame, schema::v1::RequestId, + Channel, ConnectionDriver, 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}; - use futures::future::BoxFuture; use serde_json::json; use tokio::{ net::TcpListener, @@ -259,15 +258,10 @@ mod tests { } impl AgentFactory for CapturingAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { let Channel { rx: mut incoming, tx: outgoing, @@ -301,15 +295,10 @@ mod tests { } impl AgentFactory for BatchAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (mut agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { panic!("expected one batch frame"); }; @@ -345,15 +334,10 @@ mod tests { } impl AgentFactory for FinalFrameThenExitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); let emit = self.emit.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { emit.notified().await; agent .tx @@ -377,15 +361,10 @@ mod tests { } impl AgentFactory for FinalFrameAfterInputCloseAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, ConnectionDriver) { let (agent, transport) = Channel::duplex(); let emit = self.emit.clone(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { drop(agent.rx); emit.notified().await; agent 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 7a3c1c68..65eb5317 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,7 +1,7 @@ //! HTTP-based MCP bridge transport. use agent_client_protocol::{ - BoxFuture, Channel, ConnectTo, RawJsonRpcMessage, RawJsonRpcParams, + Channel, ConnectTo, ConnectionDriver, RawJsonRpcMessage, RawJsonRpcParams, RawJsonRpcResponse as RpcResponse, TransportBatchEntry, TransportFrame, role::mcp, schema::v1::{Notification as RpcNotification, Request as RpcRequest, RequestId}, @@ -74,17 +74,15 @@ impl ConnectTo for HttpMcpBridge { } } - fn into_channel_and_future( - self, - ) -> ( - Channel, - BoxFuture<'static, Result<(), agent_client_protocol::Error>>, - ) + fn into_channel_and_future(self) -> (Channel, ConnectionDriver) where Self: Sized, { let (channel_a, channel_b) = Channel::duplex(); - (channel_a, Box::pin(run(self.listener, channel_b))) + ( + channel_a, + ConnectionDriver::new(run(self.listener, channel_b)), + ) } } diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index 8412af36..d9d537f2 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -9,9 +9,16 @@ error extension fields without ACP interpretation. Raw adapters must use the new response type. Typed ACP consumers still receive `Error`; conversion to ACP is explicit through `RawJsonRpcError::into_acp_error`. +- **Breaking:** `ConnectTo::into_channel_and_future` now returns + `(Channel, ConnectionDriver)` instead of a boxed future. Custom overrides + must distinguish owned connection work from passive endpoints; wrappers + should preserve the returned driver rather than erase its lifecycle metadata. + See the [connection-driver migration guide](../../md/migration-connection-drivers.md). ### Added +- Add `ConnectionDriver::new`, `passive`, and `is_passive`. Drivers remain + awaitable; passive readiness does not signal transport EOF. - Add a default-enabled `schemars` feature that forwards JSON Schema support to the schema crate and gates the typed MCP tool helpers. Set `default-features = false` to use the core SDK without `schemars`; custom MCP @@ -19,6 +26,14 @@ Existing users of `default-features = false` who need the previous JSON Schema or typed MCP tool APIs should add `features = ["schemars"]`. +### Fixed + +- Preserve both directions of passive channel bridges after a write + half-close, allowing a final reverse-direction response. +- Drain accepted output when an owned component finishes, including through + line and byte-stream adapters, without requiring unrelated remote input to + close. Ready component and I/O failures remain authoritative during drain. + ## [2.2.0](https://github.com/agentclientprotocol/rust-sdk/compare/v2.1.0...v2.2.0) - 2026-09-18 ### Added diff --git a/src/agent-client-protocol/src/acp_agent.rs b/src/agent-client-protocol/src/acp_agent.rs index db5dc7ae..ba0cca2a 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 { diff --git a/src/agent-client-protocol/src/component.rs b/src/agent-client-protocol/src/component.rs index ba8b9297..fd03e737 100644 --- a/src/agent-client-protocol/src/component.rs +++ b/src/agent-client-protocol/src/component.rs @@ -27,10 +27,96 @@ //! ``` 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}; +/// Drives an endpoint and declares who owns its connection lifetime. +/// +/// An active driver owns the endpoint: successful completion means no further +/// output is expected, and adapters must drain output already accepted before +/// terminating. An error terminates the connection immediately. +/// +/// A passive driver has no work to drive. Its readiness does **not** mean EOF: +/// the channel's two halves independently determine the endpoint's lifetime. +/// Adapters must inspect [`Self::is_passive`] before using completion as a +/// shutdown signal, and poll active drivers concurrently with channel traffic. +#[must_use = "active connection drivers must be polled to make progress"] +pub struct ConnectionDriver { + future: Option>>, + finish: Option>, +} + +impl ConnectionDriver { + /// Create an active driver that owns the endpoint's lifetime. + /// + /// Custom normalized adapters must finish accepted output when their channel + /// input closes. This constructor cannot externally flush or shut down an + /// opaque future that waits for additional, independently owned input. + pub fn new(future: impl Future> + Send + 'static) -> Self { + Self { + future: Some(Box::pin(future)), + finish: None, + } + } + + /// Create a passive driver for an endpoint whose channel halves own its lifetime. + pub fn passive() -> Self { + Self { + future: None, + finish: None, + } + } + + /// Whether this driver is passive, so readiness must not be treated as EOF. + #[must_use] + pub fn is_passive(&self) -> bool { + self.future.is_none() + } + + // Physical transports can finish their write half without waiting for read + // EOF. Keep this coordination private; arbitrary futures cannot support it. + pub(crate) fn with_finish( + future: impl Future> + Send + 'static, + finish: futures::channel::oneshot::Sender<()>, + ) -> Self { + Self { + future: Some(Box::pin(future)), + finish: Some(finish), + } + } + + pub(crate) fn take_finish(&mut self) -> Option> { + self.finish.take() + } +} + +impl Future for ConnectionDriver { + type Output = Result<()>; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + match &mut self.future { + 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()) + .field("finishable", &self.finish.is_some()) + .finish_non_exhaustive() + } +} + /// 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`)). /// @@ -129,7 +215,7 @@ pub trait ConnectTo: Send + 'static { client: impl ConnectTo, ) -> impl Future> + Send; - /// Convert this component into a channel endpoint and connection future. + /// Convert this component into a channel endpoint and connection driver. /// /// The returned [`Channel`] is the canonical frame-aware boundary. It carries /// complete [`TransportFrame`](crate::TransportFrame) values so default @@ -137,7 +223,7 @@ 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 /// /// 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 +232,16 @@ pub trait ConnectTo: Send + 'static { /// /// # Returns /// - /// A tuple of `(Channel, BoxFuture)` 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<()>>) + /// A tuple of `(Channel, ConnectionDriver)` where the channel is for the caller + /// to use and active drivers must be polled concurrently with channel traffic. + /// Successful active completion ends the endpoint after draining accepted + /// output; passive readiness is not EOF and preserves both channel half-closes. + 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 +259,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 +282,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 +336,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() } } @@ -269,6 +354,30 @@ mod tests { use super::*; use crate::role::UntypedRole; + #[test] + fn passive_readiness_does_not_claim_an_owned_lifetime() { + let driver = ConnectionDriver::passive(); + assert!(driver.is_passive()); + futures::executor::block_on(driver).unwrap(); + } + + #[test] + fn active_driver_preserves_errors_and_polls_unpinned() { + let error = crate::Error::internal_error().data("driver failure"); + let mut driver = ConnectionDriver::new(futures::future::ready(Err(error.clone()))); + assert!(!driver.is_passive()); + assert_eq!(futures::executor::block_on(&mut driver), Err(error)); + // Completion does not reclassify an owned endpoint as passive. + assert!(!driver.is_passive()); + } + + #[test] + fn type_erasure_preserves_passive_lifetime() { + let (channel, _other) = Channel::duplex(); + let (_, driver) = DynConnectTo::::new(channel).into_channel_and_future(); + assert!(driver.is_passive()); + } + #[test] fn dyn_connect_to_reports_static_type_name_and_correct_debug_label() { let (channel, _other) = Channel::duplex(); diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 26bed388..dcc175d2 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -6266,52 +6266,39 @@ where Self { outgoing, incoming } } - fn into_channel_transport(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { + fn into_channel_transport(self) -> (Channel, crate::ConnectionDriver) { let Self { outgoing, incoming } = self; let (channel_for_caller, channel_for_lines) = Channel::duplex(); - - let server_future = Box::pin(async move { - let Channel { rx, tx } = channel_for_lines; - let outgoing_future = transport_actor::transport_outgoing_lines_actor(rx, outgoing); - let incoming_future = transport_actor::transport_incoming_lines_actor(incoming, tx); - futures::try_join!(outgoing_future, incoming_future)?; - Ok(()) - }); - - (channel_for_caller, server_future) - } -} - -impl ConnectTo for Lines -where - OutgoingSink: futures::Sink + Send + 'static, - IncomingStream: futures::Stream> + Send + 'static, -{ - async fn connect_to(self, client: impl ConnectTo) -> Result<(), crate::Error> { - let Self { outgoing, incoming } = self; - let (Channel { rx, tx }, client_channel) = Channel::duplex(); - let close_client_output = client_channel.tx.clone(); - let client_future = Box::pin(async move { - let result = client.connect_to(client_channel).await; - close_client_output.close_channel(); - result + let Channel { mut rx, tx } = channel_for_lines; + let (finish_tx, finish_rx) = oneshot::channel(); + let finish = async move { + // Losing a finish handle is not a shutdown request. + if finish_rx.await.is_err() { + future::pending::<()>().await; + } + } + .boxed() + .shared(); + let outgoing_frames = futures::stream::poll_fn({ + let mut finish = finish.clone(); + let mut finishing = false; + move |cx| { + if !finishing && std::pin::Pin::new(&mut finish).poll(cx).is_ready() { + rx.close(); + finishing = true; + } + rx.poll_next_unpin(cx) + } }); - - // Once the client completes successfully, its incoming channel is - // gone. Keep consuming successful messages from the physical read - // half without forwarding them so a full-duplex peer cannot block our - // outgoing sink while it is being drained. Transport errors must still - // fail the connection. let discard_incoming = Arc::new(AtomicBool::new(false)); let incoming = incoming.filter_map({ let discard_incoming = discard_incoming.clone(); move |item| { - let discard_incoming = discard_incoming.load(Ordering::Acquire); - future::ready((!discard_incoming || item.is_err()).then_some(item)) + let discard = discard_incoming.load(Ordering::Acquire); + future::ready((!discard || item.is_err()).then_some(item)) } }); - - let outgoing = transport_actor::transport_outgoing_lines_actor(rx, outgoing) + let outgoing = transport_actor::transport_outgoing_lines_actor(outgoing_frames, outgoing) .boxed() .shared(); let serve_self = Box::pin({ @@ -6324,30 +6311,52 @@ where Ok(()) } }); + let server_future = crate::ConnectionDriver::with_finish( + async move { + match future::select(finish, serve_self).await { + Either::Left(((), serve_self)) => { + discard_incoming.store(true, Ordering::Release); + // Keep reading while flushing, but do not require remote + // read EOF. Poll incoming errors before clean sink drain. + match future::select(serve_self, outgoing).await { + Either::Left((result, _)) | Either::Right((result, _)) => result, + } + } + Either::Right((result, _)) => result, + } + }, + finish_tx, + ); + + (channel_for_caller, server_future) + } +} + +impl ConnectTo for Lines +where + OutgoingSink: futures::Sink + Send + 'static, + IncomingStream: futures::Stream> + Send + 'static, +{ + async fn connect_to(self, client: impl ConnectTo) -> Result<(), crate::Error> { + let (channel, mut serve_self) = self.into_channel_transport(); + let finish = serve_self + .take_finish() + .expect("built-in Lines transport supports explicit finishing"); + let client_future = Box::pin(ConnectTo::::connect_to(channel, client)); match futures::future::select(client_future, serve_self).await { Either::Left((result, serve_self)) => { result?; - discard_incoming.store(true, Ordering::Release); - - // Drive the read half while waiting for the write half, but do - // not require the peer's independent incoming stream to reach - // EOF. If incoming processing finishes successfully first, - // the shared outgoing future still owns and drains the sink. - // A successful `serve_self` result includes its shared - // outgoing clone, while any error must remain authoritative - // instead of being hidden behind the other handle. Poll it - // first so a ready read error wins over clean outgoing - // completion. - match future::select(serve_self, outgoing).await { - Either::Left((result, _)) | Either::Right((result, _)) => result, - } + // The local bridge has transferred all accepted client output. + // Finish the physical sink without waiting for remote read EOF. + let _ = finish.send(()); + serve_self.await } Either::Right((result, _)) => result, } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { + fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { self.into_channel_transport() } } @@ -6448,7 +6457,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()) } } @@ -6509,6 +6518,55 @@ impl Channel { Ok(()) } + /// Copy output concurrently with its owning driver, then drain accepted frames. + /// Passive endpoints instead retain the channel's independent half-close lifetime. + pub(crate) async fn copy_with_driver( + mut self, + mut driver: crate::ConnectionDriver, + ) -> Result<(), crate::Error> { + if driver.is_passive() { + return self.copy().await; + } + + let mut done = false; + loop { + let frame = if done { + self.rx.next().await + } else { + // Driver errors remain authoritative even when output EOF is ready. + let event = future::poll_fn(|cx| { + if let std::task::Poll::Ready(result) = std::pin::Pin::new(&mut driver).poll(cx) + { + return std::task::Poll::Ready(Either::Left(result)); + } + self.rx.poll_next_unpin(cx).map(Either::Right) + }) + .await; + match event { + Either::Left(result) => { + result?; + done = true; + self.rx.close(); + continue; + } + Either::Right(frame) => frame, + } + }; + let Some(frame) = frame else { + break; + }; + self.tx + .unbounded_send(frame) + .map_err(crate::util::internal_error)?; + } + // Propagate this half-close before waiting for a still-running driver. + drop(self); + if !done { + driver.await?; + } + Ok(()) + } + /// Bridge two endpoints while inspecting every valid message. /// /// Observers are invoked in source order, including for each valid member of @@ -6561,24 +6619,37 @@ 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 ((), (), ()) = futures::try_join!( + let passive = client_future.is_passive(); + let outgoing = Box::pin( Channel { rx: client_channel.rx, tx: self.tx, } - .copy(), + .copy_with_driver(client_future), + ); + let incoming = Box::pin( Channel { rx: self.rx, tx: client_channel.tx, } .copy(), - client_future, - )?; - Ok(()) + ); + if passive { + futures::try_join!(outgoing, incoming)?; + return Ok(()); + } + + match future::select(outgoing, incoming).await { + Either::Left((result, _)) => result, + Either::Right((result, outgoing)) => { + result?; + outgoing.await + } + } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { - (self, Box::pin(future::ready(Ok(())))) + fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { + (self, crate::ConnectionDriver::passive()) } } @@ -6586,6 +6657,53 @@ impl ConnectTo for Channel { mod tests { use super::*; + #[test] + fn dropping_unused_finish_signal_preserves_physical_half_closes() { + let outgoing = futures::sink::unfold((), |(), _line: String| { + future::ready(Ok::<_, std::io::Error>(())) + }); + let (incoming_tx, incoming_rx) = mpsc::unbounded(); + let (Channel { mut rx, tx }, mut driver) = + Lines::new(outgoing, incoming_rx).into_channel_transport(); + + drop( + driver + .take_finish() + .expect("built-in Lines driver is finishable"), + ); + drop(tx); + assert!((&mut driver).now_or_never().is_none()); + incoming_tx + .unbounded_send(Ok( + r#"{"jsonrpc":"2.0","method":"test/after-output-eof"}"#.into() + )) + .unwrap(); + assert!((&mut driver).now_or_never().is_none()); + assert!(rx.next().now_or_never().unwrap().is_some()); + + drop(incoming_tx); + futures::executor::block_on(driver).unwrap(); + assert!(rx.next().now_or_never().unwrap().is_none()); + } + + #[test] + fn explicit_physical_finish_does_not_hide_a_ready_read_error() { + let outgoing = futures::sink::unfold((), |(), _line: String| { + future::ready(Ok::<_, std::io::Error>(())) + }); + let incoming = futures::stream::iter([Err(std::io::Error::other("finish read failed"))]); + let (_channel, mut driver) = Lines::new(outgoing, incoming).into_channel_transport(); + driver.take_finish().unwrap().send(()).unwrap(); + + let error = futures::executor::block_on(driver).unwrap_err(); + assert_eq!( + error + .data + .and_then(|value| value.as_str().map(str::to_owned)), + Some("finish read failed".into()) + ); + } + #[cfg(feature = "unstable_protocol_v2")] fn connection_with_task_receiver() -> ( ConnectionTo, diff --git a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs index 62ba8f08..fcbb3664 100644 --- a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs @@ -202,7 +202,7 @@ fn malformed_line_value(raw: String) -> Result { } pub(super) async fn transport_outgoing_lines_actor( - transport_rx: mpsc::UnboundedReceiver, + transport_rx: impl futures::Stream, outgoing_lines: impl futures::Sink, ) -> Result<(), crate::Error> { transport_outgoing_frames_actor(transport_rx, outgoing_lines).await diff --git a/src/agent-client-protocol/src/lib.rs b/src/agent-client-protocol/src/lib.rs index b60787b7..6165cfec 100644 --- a/src/agent-client-protocol/src/lib.rs +++ b/src/agent-client-protocol/src/lib.rs @@ -162,7 +162,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/role/acp.rs b/src/agent-client-protocol/src/role/acp.rs index f1d54523..7892aa24 100644 --- a/src/agent-client-protocol/src/role/acp.rs +++ b/src/agent-client-protocol/src/role/acp.rs @@ -957,7 +957,8 @@ async fn reject_initialize( frame: &TransportFrame, error: crate::Error, ) -> Result<(), crate::Error> { - let RunningProtocolPeer { mut rx, tx, future } = client; + let RunningProtocolPeer { mut rx, tx, driver } = client; + let future = driver.into_driver(); send_initialize_error(&tx, frame, error)?; drop(tx); @@ -969,37 +970,89 @@ async fn reject_initialize( Ok::<_, crate::Error>(()) }; - let ((), ()) = futures::try_join!(future, drain_incoming)?; - Ok(()) + if future.is_passive() { + return drain_incoming.await; + } + + match future::select(future, Box::pin(drain_incoming)).await { + future::Either::Left((result, _)) => result, + future::Either::Right((result, future)) => { + result?; + future.await + } + } } #[cfg(feature = "unstable_protocol_v2")] struct RunningProtocolPeer { rx: futures::channel::mpsc::UnboundedReceiver, tx: futures::channel::mpsc::UnboundedSender, - future: crate::BoxFuture<'static, Result<(), crate::Error>>, + driver: ProtocolPeerDriver, +} + +#[cfg(feature = "unstable_protocol_v2")] +enum ProtocolPeerDriver { + Passive, + Active(crate::ConnectionDriver), + Completed, +} + +#[cfg(feature = "unstable_protocol_v2")] +impl ProtocolPeerDriver { + fn into_driver(self) -> crate::ConnectionDriver { + match self { + Self::Passive => crate::ConnectionDriver::passive(), + Self::Active(driver) => driver, + // Conversion happens only when handing the peer to its final + // bridge, never while reading its remaining queued frames. + Self::Completed => crate::ConnectionDriver::new(future::ready(Ok(()))), + } + } } #[cfg(feature = "unstable_protocol_v2")] impl RunningProtocolPeer { fn new(component: impl ConnectTo) -> Self { let (Channel { rx, tx }, future) = component.into_channel_and_future(); - Self { rx, tx, future } + let driver = if future.is_passive() { + ProtocolPeerDriver::Passive + } else { + ProtocolPeerDriver::Active(future) + }; + Self { rx, tx, driver } } 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 }))) - } - future::Either::Left((None, future)) => { + let Self { mut rx, tx, driver } = self; + let ProtocolPeerDriver::Active(future) = driver else { + return Ok(rx + .next() + .await + .map(|frame| (frame, Self { rx, tx, driver }))); + }; + + // Poll the owned driver first: a ready error must not be hidden by + // an equally ready frame or clean channel EOF. + match future::select(future, Box::pin(rx.next())).await { + future::Either::Right((Some(frame), future)) => Ok(Some(( + frame, + Self { + rx, + tx, + driver: ProtocolPeerDriver::Active(future), + }, + ))), + future::Either::Right((None, future)) => { + drop(tx); future.await?; Ok(None) } - future::Either::Right((result, next_message)) => { + future::Either::Left((result, next_message)) => { result?; drop(next_message); + // No more output may be accepted from an owned endpoint once + // its driver completes, even if a sender escaped the component. + rx.close(); let Some(frame) = rx.next().await else { return Ok(None); }; @@ -1008,7 +1061,7 @@ impl RunningProtocolPeer { Self { rx, tx, - future: Box::pin(future::ready(Ok(()))), + driver: ProtocolPeerDriver::Completed, }, ))) } @@ -1071,19 +1124,17 @@ async fn pipe_protocol_peers_until_closed( left: RunningProtocolPeer, right: RunningProtocolPeer, ) -> Result<(), crate::Error> { - let ((), (), (), ()) = futures::try_join!( - left.future, - right.future, + let ((), ()) = futures::try_join!( Channel { rx: left.rx, tx: right.tx, } - .copy(), + .copy_with_driver(left.driver.into_driver()), Channel { rx: right.rx, tx: left.tx, } - .copy(), + .copy_with_driver(right.driver.into_driver()), )?; Ok(()) @@ -1094,28 +1145,52 @@ async fn pipe_protocol_peers_until_done( left: RunningProtocolPeer, right: RunningProtocolPeer, ) -> Result<(), crate::Error> { - let bridge = Box::pin(async move { - let ((), ()) = futures::try_join!( - Channel { - rx: left.rx, - tx: right.tx, + let mut left_driver = left.driver.into_driver(); + let mut right_driver = right.driver.into_driver(); + let left_passive = left_driver.is_passive(); + let right_passive = right_driver.is_passive(); + let left_finish = left_driver.take_finish(); + let right_finish = right_driver.take_finish(); + let left_to_right = Box::pin( + Channel { + rx: left.rx, + tx: right.tx, + } + .copy_with_driver(left_driver), + ); + let right_to_left = Box::pin( + Channel { + rx: right.rx, + tx: left.tx, + } + .copy_with_driver(right_driver), + ); + + match future::select(left_to_right, right_to_left).await { + future::Either::Left((result, right_to_left)) => { + result?; + if left_passive || !right_passive { + if !left_passive && let Some(finish) = right_finish { + let _ = finish.send(()); + } + right_to_left.await + } else { + // Passive input may remain independently open, but a ready + // forwarding error must still beat foreground success. + crate::util::run_until(right_to_left, future::ready(Ok(()))).await } - .copy(), - Channel { - rx: right.rx, - tx: left.tx, + } + future::Either::Right((result, left_to_right)) => { + result?; + if right_passive || !left_passive { + if !right_passive && let Some(finish) = left_finish { + let _ = finish.send(()); + } + left_to_right.await + } else { + crate::util::run_until(left_to_right, future::ready(Ok(()))).await } - .copy(), - )?; - Ok(()) - }); - - match future::select(left.future, future::select(right.future, bridge)).await { - future::Either::Left((result, _)) - | future::Either::Right(( - future::Either::Left((result, _)) | future::Either::Right((result, _)), - _, - )) => result, + } } } @@ -1724,3 +1799,192 @@ where format!("ProxySessionMessages({})", self.session_id) } } + +#[cfg(all(test, feature = "unstable_protocol_v2"))] +mod lifetime_tests { + use super::*; + use crate::{ConnectionDriver, UntypedRole}; + use futures::FutureExt as _; + + fn frame() -> TransportFrame { + TransportFrame::parse_json(r#"{"jsonrpc":"2.0","method":"test/queued","params":{}}"#) + } + + #[tokio::test] + async fn passive_initialization_waits_for_a_frame_not_driver_readiness() { + let (channel, remote) = Channel::duplex(); + let peer = RunningProtocolPeer::new::(channel); + let mut next = Box::pin(peer.next_frame()); + assert!(next.as_mut().now_or_never().is_none()); + + remote.tx.unbounded_send(frame()).unwrap(); + let (_, peer) = next.await.unwrap().expect("passive peer remains connected"); + assert!(matches!(peer.driver, ProtocolPeerDriver::Passive)); + } + + #[tokio::test] + async fn active_initialization_drains_accepted_frames_without_escaped_sender_eof() { + let (Channel { rx, tx }, remote) = Channel::duplex(); + remote.tx.unbounded_send(frame()).unwrap(); + remote.tx.unbounded_send(frame()).unwrap(); + let mut polls = 0; + let driver = ConnectionDriver::new(future::poll_fn(move |_| { + polls += 1; + assert_eq!(polls, 1, "the completed driver must never be re-polled"); + std::task::Poll::Ready(Ok(())) + })); + let peer = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(driver), + }; + + let (_, peer) = peer.next_frame().await.unwrap().unwrap(); + assert!( + remote.tx.unbounded_send(frame()).is_err(), + "active completion must reject new output from escaped handles" + ); + let (_, peer) = peer.next_frame().await.unwrap().unwrap(); + assert!(peer.next_frame().await.unwrap().is_none()); + } + + #[tokio::test] + async fn ready_initialization_driver_error_beats_queued_frames() { + let (Channel { rx, tx }, remote) = Channel::duplex(); + remote.tx.unbounded_send(frame()).unwrap(); + let error = crate::Error::internal_error().data("owned initialization failed"); + let peer = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Err( + error.clone() + )))), + }; + + match peer.next_frame().await { + Err(actual) => assert_eq!(actual, error), + Ok(_) => panic!("ready driver error must not be hidden by a queued frame"), + } + } + + struct QueuedFinalClient; + + impl ConnectTo for QueuedFinalClient { + async fn connect_to(self, agent: impl ConnectTo) -> Result<(), crate::Error> { + let (mut channel, driver) = agent.into_channel_and_future(); + let foreground = async move { + channel + .tx + .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + "initialize".into(), + serde_json::json!({ "protocolVersion": 1, "clientCapabilities": {} }), + RequestId::Number(1), + )?)) + .unwrap(); + assert!(channel.rx.next().await.is_some(), "initialize response"); + // Exceed the physical writer capacity so only concurrent polling + // of the sink can make this bridge finish. + for index in 0..3 { + channel + .tx + .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::notification( + "test/final".into(), + serde_json::json!({ "index": index, "payload": "x".repeat(1024) }), + )?)) + .unwrap(); + } + Ok(()) + }; + crate::util::run_until(driver, foreground).await + } + } + + #[tokio::test] + async fn connector_completion_flushes_byte_streams_without_remote_read_eof() { + use tokio::io::{AsyncBufReadExt as _, AsyncWriteExt as _, BufReader}; + use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _}; + + let (writer, remote_reader) = tokio::io::duplex(64); + let (mut remote_input, reader) = tokio::io::duplex(64); + let physical = crate::ByteStreams::new(writer.compat_write(), reader.compat()); + let mut physical = Some(crate::DynConnectTo::::new(physical)); + let connector = tokio::spawn( + ClientProtocolConnector::new() + .with_v1(|| QueuedFinalClient) + .connect_to(move || physical.take().expect("one physical connection")), + ); + let mut lines = BufReader::new(remote_reader).lines(); + + tokio::time::timeout(std::time::Duration::from_secs(5), async { + let initialize = lines + .next_line() + .await + .unwrap() + .expect("initialize request"); + let value: serde_json::Value = serde_json::from_str(&initialize).unwrap(); + assert_eq!(value["method"], "initialize"); + remote_input + .write_all(b"{\"jsonrpc\":\"2.0\",\"id\":1,\"result\":{\"protocolVersion\":1}}\n") + .await + .unwrap(); + for index in 0..3 { + let line = lines + .next_line() + .await + .unwrap() + .expect("accepted final output"); + let value: serde_json::Value = serde_json::from_str(&line).unwrap(); + assert_eq!(value["params"]["index"], index); + assert_eq!(value["params"]["payload"].as_str().unwrap().len(), 1024); + } + connector.await.unwrap().unwrap(); + assert!(lines.next_line().await.unwrap().is_none()); + }) + .await + .expect("physical flush must not wait for independent remote input EOF"); + drop(remote_input); + } + + #[tokio::test] + async fn foreground_completion_does_not_hide_opposed_ready_driver_error() { + let (Channel { rx, tx }, _foreground_remote) = Channel::duplex(); + let foreground = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))), + }; + let (Channel { rx, tx }, _opposed_remote) = Channel::duplex(); + let error = crate::Error::internal_error().data("opposed driver failed"); + let opposed = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Err( + error.clone() + )))), + }; + + assert_eq!( + pipe_protocol_peers_until_done(foreground, opposed).await, + Err(error) + ); + } + + #[tokio::test] + async fn passive_protocol_bridge_preserves_the_reverse_half_after_eof() { + let (left, mut remote_left) = Channel::duplex(); + let (right, mut remote_right) = Channel::duplex(); + let mut bridge = Box::pin(pipe_protocol_peers_until_done( + RunningProtocolPeer::new::(left), + RunningProtocolPeer::new::(right), + )); + remote_left.tx.close_channel(); + assert!(bridge.as_mut().now_or_never().is_none()); + assert!(remote_right.rx.next().await.is_none()); + + remote_right.tx.unbounded_send(frame()).unwrap(); + remote_right.tx.close_channel(); + bridge.await.unwrap(); + assert!(remote_left.rx.next().await.is_some()); + assert!(remote_left.rx.next().await.is_none()); + } +} diff --git a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs index 8849b8a3..e56f778e 100644 --- a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs +++ b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs @@ -12,9 +12,9 @@ use std::{ }; use agent_client_protocol::{ - ByteStreams, Channel, ConnectTo, ConnectionTo, Dispatch, Error, Handled, JsonRpcMessage, - JsonRpcRequest, Lines, RawJsonRpcMessage, RawJsonRpcResponse as Response, TransportFrame, - UntypedMessage, is_incoming_transport_closed, + ByteStreams, Channel, ConnectTo, ConnectionDriver, ConnectionTo, Dispatch, DynConnectTo, Error, + Handled, JsonRpcMessage, JsonRpcRequest, Lines, RawJsonRpcMessage, + RawJsonRpcResponse as Response, TransportFrame, UntypedMessage, is_incoming_transport_closed, role::{Role, UntypedRole}, schema::v1::RequestId, }; @@ -169,6 +169,75 @@ impl ConnectTo for ImmediateClient { } } +struct CompletingClient(Result<(), Error>); + +impl ConnectTo for CompletingClient { + async fn connect_to(self, transport: impl ConnectTo) -> Result<(), Error> { + let (channel, driver) = transport.into_channel_and_future(); + for sequence in 0..3 { + channel + .tx + .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::notification( + "completed".into(), + serde_json::json!({ "sequence": sequence }), + )?)) + .map_err(Error::into_internal_error)?; + } + drop(channel); + driver.await?; + self.0 + } +} + +struct RequestReplyClient; + +impl ConnectTo for RequestReplyClient { + async fn connect_to(self, transport: impl ConnectTo) -> Result<(), Error> { + let (mut channel, driver) = transport.into_channel_and_future(); + channel + .tx + .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + "myRequest".into(), + serde_json::json!({}), + RequestId::Number(42), + )?)) + .map_err(Error::into_internal_error)?; + + // Completion depends on output being copied while this component is + // still running, not only after its driver has completed. + let Some(TransportFrame::Single(RawJsonRpcMessage::Response(Response::Result { + id, + result, + }))) = channel.rx.next().await + else { + panic!("active component lost its response"); + }; + assert_eq!(id, RequestId::Number(42)); + assert_eq!(result, serde_json::json!({ "status": "received" })); + drop(channel); + driver.await + } +} + +struct DrivenEndpoint { + channel: Channel, + driver: ConnectionDriver, +} + +impl ConnectTo for DrivenEndpoint { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + futures::try_join!( + ConnectTo::::connect_to(self.channel, client), + self.driver, + )?; + Ok(()) + } + + fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { + (self.channel, self.driver) + } +} + fn assert_connection_closed(error: &Error, method: &str) { assert!(is_incoming_transport_closed(error)); assert_eq!(error.message, "Incoming transport closed"); @@ -277,6 +346,238 @@ async fn assert_connect_to_flushes_final_response( .expect("clean EOF should succeed after flushing the response"); } +async fn assert_passive_bridge_preserves_half_close( + client: impl ConnectTo, + mut client_peer: Channel, +) { + let (transport, mut transport_peer) = Channel::duplex(); + let bridge = ConnectTo::::connect_to(transport, client); + tokio::pin!(bridge); + + // Poll with both directions open and idle. A passive driver's Ready(Ok) + // must not be mistaken for endpoint completion. + assert!( + bridge.as_mut().now_or_never().is_none(), + "passive readiness ended an open bridge" + ); + transport_peer + .tx + .unbounded_send(TransportFrame::Single( + RawJsonRpcMessage::request( + "myRequest".into(), + serde_json::json!({}), + RequestId::Number(43), + ) + .unwrap(), + )) + .unwrap(); + transport_peer.tx.close_channel(); + + tokio::time::timeout(TIMEOUT, async { + let frame = tokio::select! { + biased; + result = &mut bridge => panic!("bridge ended before forwarding the request: {result:?}"), + frame = client_peer.rx.next() => frame, + }; + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = frame else { + panic!("passive bridge lost the request"); + }; + assert_eq!(request.id, RequestId::Number(43)); + assert_eq!(&*request.method, "myRequest"); + + let eof = tokio::select! { + biased; + result = &mut bridge => panic!("one half-close ended the bridge: {result:?}"), + frame = client_peer.rx.next() => frame, + }; + assert!(eof.is_none(), "write half-close was not forwarded"); + assert!( + bridge.as_mut().now_or_never().is_none(), + "bridge must wait for its still-open reverse direction" + ); + + // Only send the final response after observing EOF in the request + // direction, so buffering cannot conceal premature bridge shutdown. + client_peer + .tx + .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + RequestId::Number(43), + Ok(serde_json::json!({ "status": "received" })), + ))) + .expect("reverse direction must remain open after the first EOF"); + client_peer.tx.close_channel(); + + let receive_response = async { + let Some(TransportFrame::Single(RawJsonRpcMessage::Response(Response::Result { + id, + result, + }))) = transport_peer.rx.next().await + else { + panic!("passive bridge lost the final reverse-direction response"); + }; + assert_eq!(id, RequestId::Number(43)); + assert_eq!(result, serde_json::json!({ "status": "received" })); + assert!( + transport_peer.rx.next().await.is_none(), + "second half-close was not forwarded" + ); + }; + let (result, ()) = join(&mut bridge, receive_response).await; + result.expect("both closed halves should complete the passive bridge cleanly"); + }) + .await + .expect("passive bridge hung while forwarding half-closes"); +} + +#[test] +fn channel_and_erased_channel_drivers_are_explicitly_passive() { + for erased in [false, true] { + let (endpoint, _peer) = Channel::duplex(); + let (_channel, driver) = if erased { + DynConnectTo::::new(endpoint).into_channel_and_future() + } else { + ConnectTo::::into_channel_and_future(endpoint) + }; + assert!(driver.is_passive()); + driver + .now_or_never() + .expect("passive driver should be immediately ready") + .expect("passive readiness should succeed"); + } + + let driver = ConnectionDriver::new(future::ready(Ok(()))); + assert!( + !driver.is_passive(), + "ready owned work is active even when it completes immediately" + ); + driver.now_or_never().unwrap().unwrap(); +} + +#[tokio::test] +async fn passive_channel_bridge_waits_for_both_halves() { + let (client, client_peer) = Channel::duplex(); + assert_passive_bridge_preserves_half_close(client, client_peer).await; +} + +#[tokio::test] +async fn erased_passive_channel_bridge_waits_for_both_halves() { + let (client, client_peer) = Channel::duplex(); + assert_passive_bridge_preserves_half_close( + DynConnectTo::::new(client), + client_peer, + ) + .await; +} + +#[tokio::test] +async fn active_channel_completion_drains_output_without_remote_eof() { + let (transport, mut peer) = Channel::duplex(); + tokio::time::timeout( + TIMEOUT, + ConnectTo::::connect_to(transport, CompletingClient(Ok(()))), + ) + .await + .expect("active completion waited for the unrelated remote sender") + .expect("active completion should drain accepted output"); + + // The bridge has already returned. Every accepted frame must be available + // now, in order, even though the remote sender was held open throughout. + for sequence in 0..3 { + let Some(TransportFrame::Single(RawJsonRpcMessage::Notification(notification))) = peer + .rx + .next() + .now_or_never() + .expect("accepted output was not drained before completion") + else { + panic!("active completion lost an accepted notification"); + }; + assert_eq!(&*notification.method, "completed"); + assert_eq!( + serde_json::to_value(notification.params).unwrap(), + serde_json::json!({ "sequence": sequence }) + ); + } + assert!(matches!(peer.rx.next().now_or_never(), Some(None))); + assert!( + peer.tx.is_closed(), + "completed bridge retained the unrelated remote input" + ); +} + +#[tokio::test] +async fn active_channel_output_is_copied_before_component_completion() { + let (transport, mut peer) = Channel::duplex(); + let connection = ConnectTo::::connect_to(transport, RequestReplyClient); + let respond = async move { + let Some(TransportFrame::Single(RawJsonRpcMessage::Request(request))) = + peer.rx.next().await + else { + panic!("bridge did not copy the active component's request"); + }; + assert_eq!(request.id, RequestId::Number(42)); + assert_eq!(&*request.method, "myRequest"); + peer.tx + .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::response( + request.id, + Ok(serde_json::json!({ "status": "received" })), + ))) + .unwrap(); + peer + }; + + let peer = tokio::time::timeout(TIMEOUT, async { + let (result, peer) = join(connection, respond).await; + result.expect("request/reply component should complete successfully"); + peer + }) + .await + .expect("outbound copying waited for a component that needed its reply first"); + assert!(peer.tx.is_closed()); +} + +#[tokio::test] +async fn clean_channel_drain_does_not_hide_ready_driver_error() { + let (transport, _remote) = Channel::duplex(); + let (channel, component_peer) = Channel::duplex(); + drop(component_peer); + let error = tokio::time::timeout( + TIMEOUT, + ConnectTo::::connect_to( + transport, + DrivenEndpoint { + channel, + driver: ConnectionDriver::new(future::ready(Err( + Error::internal_error().data("ready driver failed") + ))), + }, + ), + ) + .await + .expect("ready driver error waited for remote EOF") + .expect_err("successful output drain must not mask an active driver error"); + assert_eq!(error.data, Some(serde_json::json!("ready driver failed"))); +} + +#[tokio::test] +async fn clean_lines_drain_does_not_hide_ready_component_error() { + let outgoing = futures::sink::unfold((), |(), _line: String| async { Ok::<_, io::Error>(()) }); + let incoming = stream::empty::>(); + let error = tokio::time::timeout( + TIMEOUT, + ConnectTo::::connect_to( + Lines::new(outgoing, incoming), + CompletingClient(Err(Error::internal_error().data("ready component failed"))), + ), + ) + .await + .expect("ready component error should resolve alongside the clean transport") + .expect_err("clean transport completion must not mask the component error"); + assert_eq!( + error.data, + Some(serde_json::json!("ready component failed")) + ); +} + #[tokio::test] async fn connect_to_returns_cleanly_on_incoming_eof_with_spawned_work() { let (transport, peer) = Channel::duplex(); From fa695f44e5ae86244af0929eb776a1f74899d745 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Fri, 2 Oct 2026 12:32:27 +0200 Subject: [PATCH 2/3] refactor(acp)!: represent connection work as optional owned drivers Return None for endpoints without owned work and Some(ConnectionDriver) for genuine futures. Remove passive readiness and is_passive, keeping private finish coordination only on owned drivers. Propagate optional ownership through consumers and preserve all lifetime guarantees. Add compile-fail and absent-work half-close regressions and revise migration guidance. BREAKING CHANGE: ConnectTo::into_channel_and_future now returns (Channel, Option). Low-level callers must handle absence explicitly instead of awaiting a no-op passive driver. --- md/migration-connection-drivers.md | 44 +++-- md/transport-architecture.md | 29 ++-- .../src/snoop.rs | 80 ++++++++- src/agent-client-protocol-http/CHANGELOG.md | 8 +- src/agent-client-protocol-http/src/client.rs | 17 +- .../src/connection.rs | 106 ++++++++---- .../src/http_server.rs | 20 +-- .../src/websocket_server.rs | 16 +- .../src/mcp_over_acp/http.rs | 5 +- src/agent-client-protocol/CHANGELOG.md | 11 +- src/agent-client-protocol/src/acp_agent.rs | 9 +- src/agent-client-protocol/src/component.rs | 161 ++++++++++++------ src/agent-client-protocol/src/jsonrpc.rs | 39 +++-- src/agent-client-protocol/src/role/acp.rs | 95 ++++++++--- .../tests/jsonrpc_transport_close.rs | 34 ++-- .../tests/protocol_v2.rs | 16 +- .../tests/proxy_protocol_router_v2.rs | 2 +- 17 files changed, 482 insertions(+), 210 deletions(-) diff --git a/md/migration-connection-drivers.md b/md/migration-connection-drivers.md index d55d7a25..23042c2a 100644 --- a/md/migration-connection-drivers.md +++ b/md/migration-connection-drivers.md @@ -1,8 +1,9 @@ # Migrating Connection Drivers -`ConnectTo::into_channel_and_future` now returns `(Channel, ConnectionDriver)` -instead of `(Channel, BoxFuture<'static, Result<()>>)`. The same change applies -when accessing a component through `DynConnectTo`. +`ConnectTo::into_channel_and_future` now returns +`(Channel, Option)` instead of +`(Channel, BoxFuture<'static, Result<()>>)`. The same change applies when +accessing a component through `DynConnectTo`. This is a source-breaking transport-adapter change. It does not change ACP wire messages, the raw `Channel` sender/receiver types, or `unbounded_send`, and it @@ -12,46 +13,55 @@ introduces no new frame-size, queue, or task limits. If your component implements only `connect_to`, no change is needed. The default conversion still creates a channel pair and drives your component, now -returning an owned `ConnectionDriver`. +returning `Some(ConnectionDriver)`. -Existing callers that infer the returned type and await the driver continue -to work: +Low-level callers must handle the optional work explicitly. The optional value +is not a future: awaiting it directly no longer compiles. For a component that +is known to own work, extract its driver before polling it: ```rust,ignore let (channel, driver) = component.into_channel_and_future(); +let driver = driver.expect("this component owns connection work"); // Use channel while continuing to poll the driver. driver.await?; ``` +For a generic component, handle both cases: poll `Some(driver)` alongside +traffic and drain accepted output on completion; for `None`, retain the +channel's independent halves until they close. Absence of work is not EOF. +Do not replace `None` with a ready-success future in a shutdown race. + ## Custom conversion overrides Import `ConnectionDriver` from `agent_client_protocol` and change the return type. Wrap a future that owns the connection work with `ConnectionDriver::new`: ```rust,ignore -fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { +fn into_channel_and_future(self) -> (Channel, Option) { let (channel, future) = self.into_channel_transport(); - (channel, ConnectionDriver::new(future)) + (channel, Some(ConnectionDriver::new(future))) } ``` -For an endpoint whose work is driven elsewhere, return -`ConnectionDriver::passive()` instead of wrapping a ready no-op future: +For an endpoint whose work is driven elsewhere, return `None` instead of +wrapping a ready no-op future: ```rust,ignore -fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { - (self.channel, ConnectionDriver::passive()) +fn into_channel_and_future(self) -> (Channel, Option) { + (self.channel, None) } ``` -Passive drivers are awaitable and immediately succeed, but that success is -**not EOF**. Bridges must check `is_passive()` before using driver completion -as a lifetime signal. Passive bridges retain each read/write half until its +An existing `Channel` has no driver. There is no awaitable passive sentinel, +and no finish hook belongs to the `None` case. `ConnectionDriver` always holds +real owned work; some built-in owned drivers additionally support private +finish coordination. Passive bridges retain each read/write half until its own closure; an input half-close can still be followed by a final response. If a wrapper simply exposes another component's endpoint, return its original -`(channel, driver)` pair. Re-boxing that driver and wrapping it with `new` -would erase the passive distinction and any built-in finish coordination. +`(channel, optional_driver)` pair. Re-boxing an owned driver and wrapping it +with `new` would erase its built-in finish coordination. Inventing a ready +driver for `None` would also turn absence into a false completion signal. ## Completion and drain responsibilities diff --git a/md/transport-architecture.md b/md/transport-architecture.md index 08d72a19..1013846d 100644 --- a/md/transport-architecture.md +++ b/md/transport-architecture.md @@ -265,18 +265,24 @@ component to its counterpart and drives the connection until completion. `Channel` plus an explicit connection driver: ```rust,ignore -fn into_channel_and_future(self) -> (Channel, ConnectionDriver); +fn into_channel_and_future(self) -> (Channel, Option); ``` -The channel carries only `TransportFrame` wire events. The awaitable -`ConnectionDriver` distinguishes owned work from a passive endpoint: +The channel carries only `TransportFrame` wire events. The optional driver +distinguishes owned work from a passive endpoint: -| Driver | Meaning of successful completion | +| Returned work | Lifetime rule | | --- | --- | -| `ConnectionDriver::new(future)` | The component's owned work has finished. | -| `ConnectionDriver::passive()` | No work is owned here; readiness says nothing about either I/O half. | +| `Some(ConnectionDriver::new(future))` | Poll the owned work alongside traffic; successful completion ends it after accepted output is drained. | +| `None` | No work is owned here; each channel half determines its own lifetime. | -A raw `Channel` is passive. Its bridge preserves both directions independently: +A `ConnectionDriver` always contains a real future; the optional return value +cannot itself be awaited. There is no ready-successful passive driver and no +finish hook in the `None` case. This makes the ownership decision explicit +rather than requiring callers to recognize a special future. + +A raw `Channel` returns `None`. Its bridge preserves both directions +independently: one sender closing must not prevent a final response in the reverse direction. Owned completion lets a bridge stop accepting new output, drain frames already accepted, and finish without waiting for unrelated remote input to close. @@ -290,8 +296,9 @@ their explicit finish handling does not require remote read EOF. Merely wrapping an arbitrary future cannot make an opaque custom adapter drain safely. Most components implement only `connect_to`; default normalization supplies -the owned driver. Direct transports override `into_channel_and_future` to avoid -an intermediate copy. Wrappers that expose an existing endpoint should forward +`Some(driver)` containing the owned work. Direct transports override +`into_channel_and_future` to avoid an intermediate copy. Wrappers that expose an +existing endpoint should forward its driver unchanged so passive identity and built-in completion handling are not lost. @@ -372,7 +379,9 @@ Split the socket and pass compatible read/write halves to `ByteStreams::new`. Embedders supply and drive their own runtime and host transport: - Exchange `TransportFrame` values through an in-component `Channel`. A caller - using `ConnectTo::into_channel_and_future` must poll the returned future. + using `ConnectTo::into_channel_and_future` polls a present owned driver + alongside traffic. When no driver is returned, preserve the channel halves' + independent lifetimes; absence is not EOF. - Exchange newline-delimited JSON through `Lines`, using a `futures::Sink` and `futures::Stream>`. diff --git a/src/agent-client-protocol-conductor/src/snoop.rs b/src/agent-client-protocol-conductor/src/snoop.rs index bc19465a..61613873 100644 --- a/src/agent-client-protocol-conductor/src/snoop.rs +++ b/src/agent-client-protocol-conductor/src/snoop.rs @@ -45,6 +45,20 @@ impl ConnectTo for SnooperComponent { self.outgoing_message, ); + // Absence contributes no owned work to the all-half join; it is not + // an EOF signal for either direction of the channel bridge. + let client_future = async move { + if let Some(driver) = client_future { + driver.await?; + } + Ok::<(), agent_client_protocol::Error>(()) + }; + let base_future = async move { + if let Some(driver) = base_future { + driver.await?; + } + Ok::<(), agent_client_protocol::Error>(()) + }; (client_future, base_future, snoop).try_join().await?; Ok(()) } @@ -60,8 +74,11 @@ mod tests { time::Duration, }; - use agent_client_protocol::{ByteStreams, ConnectionTo, Responder, UntypedRole}; + use agent_client_protocol::{ + ByteStreams, ConnectionTo, Responder, TransportFrame, UntypedRole, + }; use agent_client_protocol_test::{MyRequest, MyResponse}; + use futures::StreamExt as _; use serde_json::{Value, json}; use tokio::io::{AsyncBufReadExt as _, AsyncWriteExt as _, BufReader}; use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _}; @@ -70,6 +87,67 @@ mod tests { const TIMEOUT: Duration = Duration::from_secs(10); + #[tokio::test(flavor = "current_thread")] + async fn absent_drivers_preserve_both_channel_halves() { + tokio::task::LocalSet::new() + .run_until(async { + let (client, mut client_peer) = Channel::duplex(); + let (base, mut base_peer) = Channel::duplex(); + let incoming_count = Arc::new(AtomicUsize::new(0)); + let outgoing_count = Arc::new(AtomicUsize::new(0)); + let observed_incoming = incoming_count.clone(); + let observed_outgoing = outgoing_count.clone(); + let snooper = SnooperComponent::::new( + base, + move |_| { + observed_incoming.fetch_add(1, Ordering::SeqCst); + Ok(()) + }, + move |_| { + observed_outgoing.fetch_add(1, Ordering::SeqCst); + Ok(()) + }, + ); + let task = tokio::task::spawn_local(snooper.connect_to(client)); + let frame = TransportFrame::Single( + RawJsonRpcMessage::notification("test/passive".into(), json!({})).unwrap(), + ); + + client_peer.tx.unbounded_send(frame.clone()).unwrap(); + assert!( + tokio::time::timeout(TIMEOUT, base_peer.rx.next()) + .await + .expect("absent drivers must not stop forwarding") + .is_some() + ); + drop(client_peer.tx); + assert!( + tokio::time::timeout(TIMEOUT, base_peer.rx.next()) + .await + .expect("client half-close must reach the base") + .is_none() + ); + assert!(!task.is_finished(), "one half-close must not end the join"); + + base_peer.tx.unbounded_send(frame).unwrap(); + assert!( + tokio::time::timeout(TIMEOUT, client_peer.rx.next()) + .await + .expect("reverse forwarding must survive the first half-close") + .is_some() + ); + assert_eq!(incoming_count.load(Ordering::SeqCst), 1); + assert_eq!(outgoing_count.load(Ordering::SeqCst), 1); + drop(base_peer.tx); + tokio::time::timeout(TIMEOUT, task) + .await + .expect("snooper must finish after both channel halves close") + .expect("snooper task panicked") + .expect("snooper connection failed"); + }) + .await; + } + #[tokio::test(flavor = "current_thread")] async fn tracing_preserves_json_rpc_batch_frames() { tokio::task::LocalSet::new() diff --git a/src/agent-client-protocol-http/CHANGELOG.md b/src/agent-client-protocol-http/CHANGELOG.md index 44c21d23..c52700bd 100644 --- a/src/agent-client-protocol-http/CHANGELOG.md +++ b/src/agent-client-protocol-http/CHANGELOG.md @@ -5,16 +5,16 @@ ### Changed - Adapt `HttpClient`'s `ConnectTo` conversion to the core SDK's breaking - `ConnectionDriver` return type. Channels and HTTP framing remain unchanged; - no new resource limits are introduced. + optional `ConnectionDriver` return type. Channels and HTTP framing remain + unchanged; no new resource limits are introduced. ### Fixed - Preserve raw JSON-RPC error codes, omitted versus null data, and error extension fields across HTTP/SSE and WebSocket transports, using the core SDK's new `RawJsonRpcResponse` representation. -- Do not treat a passive agent-factory endpoint's no-op driver as agent - completion while its transport remains open. +- Preserve HTTP channel pumps when an agent-factory endpoint has no owned + driver; absence is not agent completion while its transport remains open. - On active agent completion, reject further output from escaped sender clones and drain accepted frames before removing the connection and closing its streams, without waiting for those clones to be dropped. diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index 6af4cd73..2b4acf70 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -102,6 +102,7 @@ impl HttpClient { impl ConnectTo for HttpClient { async fn connect_to(self, client: impl ConnectTo) -> Result<(), AcpError> { let (channel, transport) = ConnectTo::::into_channel_and_future(self); + let transport = transport.expect("HttpClient owns its physical transport driver"); let shutdown_tx = channel.tx.clone(); match futures::future::select( std::pin::pin!(client.connect_to(channel)), @@ -122,9 +123,9 @@ impl ConnectTo for HttpClient { } } - fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { + fn into_channel_and_future(self) -> (Channel, Option) { let (caller, transport) = Channel::duplex(); - (caller, ConnectionDriver::new(run(self, transport))) + (caller, Some(ConnectionDriver::new(run(self, transport)))) } } @@ -1589,6 +1590,12 @@ mod tests { Ok(()) }; + let transport = async move { + if let Some(transport) = transport { + transport.await?; + } + Ok::<(), AcpError>(()) + }; let ((), ()) = futures::try_join!(transport, client)?; Ok(()) } @@ -1624,6 +1631,12 @@ mod tests { Ok(()) }; + let transport = async move { + if let Some(transport) = transport { + transport.await?; + } + Ok::<(), AcpError>(()) + }; let ((), ()) = futures::try_join!(transport, client)?; Ok(()) } diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index cd635b92..8e190c8c 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -424,7 +424,7 @@ pub(crate) struct ConnectionRegistry { } pub(crate) trait AgentFactory: Send + Sync + 'static { - fn spawn_agent(&self) -> (Channel, agent_client_protocol::ConnectionDriver); + fn spawn_agent(&self) -> (Channel, Option); } impl AgentFactory for F @@ -432,7 +432,7 @@ where F: Fn() -> C + Send + Sync + 'static, C: agent_client_protocol::ConnectTo, { - fn spawn_agent(&self) -> (Channel, agent_client_protocol::ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { self().into_channel_and_future() } } @@ -540,25 +540,27 @@ impl ConnectionRegistry { let connection_for_task = Arc::downgrade(&connection); let agent_handle = tokio::spawn(async move { let conn_id_for_agent = conn_id_for_task.clone(); - let agent = async move { - if agent_future.is_passive() { - // A passive endpoint is driven only by the channel pumps; - // its immediately ready no-op driver is not connection EOF. - std::future::pending::<()>().await; - } - if let Err(e) = agent_future.await { - error!(connection_id = %conn_id_for_agent, "ACP agent task error: {e}"); - } - }; - futures::pin_mut!(agent); - futures::pin_mut!(pump); - match futures::future::select(agent, pump).await { - futures::future::Either::Left(((), pump)) => { - inbound_abort.abort(); - let _sent = finish_outbound_tx.send(()); - pump.await; + if let Some(agent_future) = agent_future { + let agent = async move { + if let Err(e) = agent_future.await { + error!(connection_id = %conn_id_for_agent, "ACP agent task error: {e}"); + } + }; + futures::pin_mut!(agent); + futures::pin_mut!(pump); + match futures::future::select(agent, pump).await { + futures::future::Either::Left(((), pump)) => { + inbound_abort.abort(); + let _sent = finish_outbound_tx.send(()); + pump.await; + } + futures::future::Either::Right(((), _agent)) => {} } - futures::future::Either::Right(((), _agent)) => {} + } else { + // With no owned agent work, only the channel pumps determine + // completion. Do not interpret absence as connection EOF. + drop(finish_outbound_tx); + pump.await; } debug!(connection_id = %conn_id_for_task, "ACP connection task ended"); let connection_to_close = drain_connection_router(connection_for_task).await; @@ -754,7 +756,7 @@ mod tests { } impl AgentFactory for ExitingAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let exit = self.exit.clone(); let future = ConnectionDriver::new(async move { @@ -763,14 +765,14 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } struct RespondThenExitAgentFactory; impl AgentFactory for RespondThenExitAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let future = ConnectionDriver::new(async move { agent @@ -783,7 +785,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -792,7 +794,7 @@ mod tests { } impl AgentFactory for MalformedThenWaitAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let emit = self.emit.clone(); let future = ConnectionDriver::new(async move { @@ -808,7 +810,7 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } @@ -818,7 +820,7 @@ mod tests { } impl AgentFactory for SendThenWaitAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let message = self.message.clone(); let exit = self.exit.clone(); @@ -831,7 +833,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -840,7 +842,7 @@ mod tests { } impl AgentFactory for BatchThenWaitAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let exit = self.exit.clone(); let future = ConnectionDriver::new(async move { @@ -865,7 +867,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -876,7 +878,7 @@ mod tests { } impl AgentFactory for FinalFrameThenExitAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); *self.escaped_output.lock().unwrap() = Some(agent.tx.clone()); let emit = self.emit.clone(); @@ -895,12 +897,12 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } #[tokio::test] - async fn passive_agent_driver_does_not_close_http_channel_pumps() { + async fn absent_agent_driver_preserves_half_closes_and_natural_completion() { let (endpoint, mut remote) = Channel::duplex(); let endpoint = std::sync::Mutex::new(Some(endpoint)); let registry = ConnectionRegistry::new(Arc::new(move || { @@ -914,10 +916,10 @@ mod tests { assert!( timeout(Duration::from_secs(1), remote.rx.next()) .await - .expect("passive readiness must not abort inbound forwarding") + .expect("absent owned work must not abort inbound forwarding") .is_some() ); - remote.tx.unbounded_send(frame).unwrap(); + remote.tx.unbounded_send(frame.clone()).unwrap(); assert!( timeout(Duration::from_secs(1), connection.recv_initial()) .await @@ -925,8 +927,38 @@ mod tests { .is_some() ); assert!(registry.get(&connection_id).await.is_some()); - connection.shutdown().await; - registry.remove(&connection_id).await; + assert!(!*connection.subscribe_closed().borrow()); + + drop(remote.tx); + assert!( + timeout(Duration::from_secs(1), connection.recv_initial()) + .await + .expect("the outbound pump should observe its half-close") + .is_none() + ); + connection.inbound_tx.send(frame.clone()).unwrap(); + assert!( + timeout(Duration::from_secs(1), remote.rx.next()) + .await + .expect("the other half must still forward after outbound EOF") + .is_some() + ); + assert!(registry.get(&connection_id).await.is_some()); + let mut closed = connection.subscribe_closed(); + assert!(!*closed.borrow()); + + drop(remote.rx); + // Forwarding observes the closed recipient on its next send; no + // explicit registry shutdown is needed to finish the two pumps. + connection.inbound_tx.send(frame).unwrap(); + timeout(Duration::from_secs(1), async { + while !*closed.borrow() { + closed.changed().await.unwrap(); + } + }) + .await + .expect("both closed halves should finish without a synthetic driver"); + assert!(registry.get(&connection_id).await.is_none()); } #[tokio::test] diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index 5cbc7479..8271186b 100644 --- a/src/agent-client-protocol-http/src/http_server.rs +++ b/src/agent-client-protocol-http/src/http_server.rs @@ -500,7 +500,7 @@ mod tests { } impl AgentFactory for CapturingAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); let future = ConnectionDriver::new(async move { @@ -519,14 +519,14 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } struct RejectingInitializeAgentFactory; impl AgentFactory for RejectingInitializeAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (mut agent, transport) = Channel::duplex(); let future = ConnectionDriver::new(async move { match agent.rx.next().await { @@ -568,14 +568,14 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } struct PendingInitializeAgentFactory; impl AgentFactory for PendingInitializeAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let future = ConnectionDriver::new(async move { let Channel { @@ -586,7 +586,7 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } @@ -595,7 +595,7 @@ mod tests { } impl AgentFactory for BatchAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (mut agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); let future = ConnectionDriver::new(async move { @@ -651,14 +651,14 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } struct SideTrafficBeforeInitializeResponseAgentFactory; impl AgentFactory for SideTrafficBeforeInitializeResponseAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (mut agent, transport) = Channel::duplex(); let future = ConnectionDriver::new(async move { let Some(TransportFrame::Batch(batch)) = agent.rx.next().await else { @@ -695,7 +695,7 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } diff --git a/src/agent-client-protocol-http/src/websocket_server.rs b/src/agent-client-protocol-http/src/websocket_server.rs index 1ab3b8ab..56ee7d67 100644 --- a/src/agent-client-protocol-http/src/websocket_server.rs +++ b/src/agent-client-protocol-http/src/websocket_server.rs @@ -258,7 +258,7 @@ mod tests { } impl AgentFactory for CapturingAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); let future = ConnectionDriver::new(async move { @@ -286,7 +286,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -295,7 +295,7 @@ mod tests { } impl AgentFactory for BatchAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (mut agent, transport) = Channel::duplex(); let forwarded = self.forwarded.clone(); let future = ConnectionDriver::new(async move { @@ -325,7 +325,7 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } @@ -334,7 +334,7 @@ mod tests { } impl AgentFactory for FinalFrameThenExitAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let emit = self.emit.clone(); let future = ConnectionDriver::new(async move { @@ -352,7 +352,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -361,7 +361,7 @@ mod tests { } impl AgentFactory for FinalFrameAfterInputCloseAgentFactory { - fn spawn_agent(&self) -> (Channel, ConnectionDriver) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); let emit = self.emit.clone(); let future = ConnectionDriver::new(async move { @@ -380,7 +380,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } 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 65eb5317..689630c1 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 @@ -69,19 +69,20 @@ impl ConnectTo for HttpMcpBridge { client: impl ConnectTo, ) -> Result<(), agent_client_protocol::Error> { let (channel, serve_self) = self.into_channel_and_future(); + let serve_self = serve_self.expect("HttpMcpBridge owns its HTTP listener driver"); match futures::future::select(pin!(client.connect_to(channel)), serve_self).await { Either::Left((result, _)) | Either::Right((result, _)) => result, } } - fn into_channel_and_future(self) -> (Channel, ConnectionDriver) + fn into_channel_and_future(self) -> (Channel, Option) where Self: Sized, { let (channel_a, channel_b) = Channel::duplex(); ( channel_a, - ConnectionDriver::new(run(self.listener, channel_b)), + Some(ConnectionDriver::new(run(self.listener, channel_b))), ) } } diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index d9d537f2..cc659f1c 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -10,15 +10,16 @@ new response type. Typed ACP consumers still receive `Error`; conversion to ACP is explicit through `RawJsonRpcError::into_acp_error`. - **Breaking:** `ConnectTo::into_channel_and_future` now returns - `(Channel, ConnectionDriver)` instead of a boxed future. Custom overrides - must distinguish owned connection work from passive endpoints; wrappers - should preserve the returned driver rather than erase its lifecycle metadata. + `(Channel, Option)` instead of a boxed future. `Some` owns + connection work; a passive endpoint returns `None`. Custom overrides and + low-level callers must handle absence explicitly; wrappers should preserve + the original optional driver rather than erase its lifecycle metadata. See the [connection-driver migration guide](../../md/migration-connection-drivers.md). ### Added -- Add `ConnectionDriver::new`, `passive`, and `is_passive`. Drivers remain - awaitable; passive readiness does not signal transport EOF. +- Add an owned-only, awaitable `ConnectionDriver`. Passive endpoints have no + driver, so absence cannot be mistaken for a successful completed future. - Add a default-enabled `schemars` feature that forwards JSON Schema support to the schema crate and gates the typed MCP tool helpers. Set `default-features = false` to use the core SDK without `schemars`; custom MCP diff --git a/src/agent-client-protocol/src/acp_agent.rs b/src/agent-client-protocol/src/acp_agent.rs index ba0cca2a..e209e70a 100644 --- a/src/agent-client-protocol/src/acp_agent.rs +++ b/src/agent-client-protocol/src/acp_agent.rs @@ -1406,7 +1406,8 @@ mod tests { let (agent, mut pid_rx) = wrapper_agent( "echo ACP_TEST_CHILD_PID=$$ >&2; exec 1>&-; sleep 30 & child=$!; wait \"$child\"", ); - let (channel, mut connection) = crate::ConnectTo::::into_channel_and_future(agent); + let (channel, connection) = crate::ConnectTo::::into_channel_and_future(agent); + let mut connection = connection.expect("AcpAgent owns its process connection"); let crate::Channel { rx: _incoming, tx: outgoing, @@ -1448,7 +1449,8 @@ mod tests { let (agent, mut pid_rx) = wrapper_agent( "sleep 30 & child=$!; echo ACP_TEST_CHILD_PID=$child >&2; wait \"$child\"", ); - let (_channel, mut connection) = crate::ConnectTo::::into_channel_and_future(agent); + let (_channel, connection) = crate::ConnectTo::::into_channel_and_future(agent); + let mut connection = connection.expect("AcpAgent owns its process connection"); let descendant_pid = reported_descendant_pid(&mut connection, &mut pid_rx).await; let mut cleanup = KillOnDrop(Some(descendant_pid)); @@ -1464,7 +1466,8 @@ mod tests { let (agent, mut pid_rx) = wrapper_agent( "sh -c 'trap \"\" HUP; exec sleep 30' >/dev/null & child=$!; echo ACP_TEST_CHILD_PID=$child >&2; exit 17", ); - let (_channel, mut connection) = crate::ConnectTo::::into_channel_and_future(agent); + let (_channel, connection) = crate::ConnectTo::::into_channel_and_future(agent); + let mut connection = connection.expect("AcpAgent owns its process connection"); let descendant_pid = reported_descendant_pid(&mut connection, &mut pid_rx).await; let mut cleanup = KillOnDrop(Some(descendant_pid)); diff --git a/src/agent-client-protocol/src/component.rs b/src/agent-client-protocol/src/component.rs index fd03e737..72373f13 100644 --- a/src/agent-client-protocol/src/component.rs +++ b/src/agent-client-protocol/src/component.rs @@ -37,49 +37,34 @@ use std::{ use crate::{Channel, Result, role::Role}; -/// Drives an endpoint and declares who owns its connection lifetime. +/// Drives owned endpoint work. /// -/// An active driver owns the endpoint: successful completion means no further +/// A driver owns the endpoint: successful completion means no further /// output is expected, and adapters must drain output already accepted before /// terminating. An error terminates the connection immediately. /// -/// A passive driver has no work to drive. Its readiness does **not** mean EOF: -/// the channel's two halves independently determine the endpoint's lifetime. -/// Adapters must inspect [`Self::is_passive`] before using completion as a -/// shutdown signal, and poll active drivers concurrently with channel traffic. -#[must_use = "active connection drivers must be polled to make progress"] +/// Poll the driver concurrently with channel traffic. Endpoints without owned +/// work return `None` from [`ConnectTo::into_channel_and_future`], not a driver: +/// their channel halves independently determine their lifetime. +#[must_use = "connection drivers must be polled to make progress"] pub struct ConnectionDriver { - future: Option>>, + future: BoxFuture<'static, Result<()>>, finish: Option>, } impl ConnectionDriver { - /// Create an active driver that owns the endpoint's lifetime. + /// Create a driver that owns the endpoint's lifetime. /// /// Custom normalized adapters must finish accepted output when their channel /// input closes. This constructor cannot externally flush or shut down an /// opaque future that waits for additional, independently owned input. pub fn new(future: impl Future> + Send + 'static) -> Self { Self { - future: Some(Box::pin(future)), + future: Box::pin(future), finish: None, } } - /// Create a passive driver for an endpoint whose channel halves own its lifetime. - pub fn passive() -> Self { - Self { - future: None, - finish: None, - } - } - - /// Whether this driver is passive, so readiness must not be treated as EOF. - #[must_use] - pub fn is_passive(&self) -> bool { - self.future.is_none() - } - // Physical transports can finish their write half without waiting for read // EOF. Keep this coordination private; arbitrary futures cannot support it. pub(crate) fn with_finish( @@ -87,7 +72,7 @@ impl ConnectionDriver { finish: futures::channel::oneshot::Sender<()>, ) -> Self { Self { - future: Some(Box::pin(future)), + future: Box::pin(future), finish: Some(finish), } } @@ -101,17 +86,13 @@ impl Future for ConnectionDriver { type Output = Result<()>; fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { - match &mut self.future { - Some(future) => future.as_mut().poll(cx), - None => Poll::Ready(Ok(())), - } + self.future.as_mut().poll(cx) } } impl Debug for ConnectionDriver { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ConnectionDriver") - .field("passive", &self.is_passive()) .field("finishable", &self.finish.is_some()) .finish_non_exhaustive() } @@ -147,7 +128,7 @@ impl Debug for ConnectionDriver { /// Components can be used in two ways: /// /// 1. **`connect_to(client)`** - Connect directly to another component (most components implement this) -/// 2. **`into_channel_and_future()`** - Obtain a channel endpoint and a future that drives the connection +/// 2. **`into_channel_and_future()`** - Obtain a channel endpoint and optional owned driver /// /// Most components only need to implement `connect_to(client)`. The /// `into_channel_and_future()` method has a default implementation that creates an intermediate @@ -215,7 +196,7 @@ pub trait ConnectTo: Send + 'static { client: impl ConnectTo, ) -> impl Future> + Send; - /// Convert this component into a channel endpoint and connection driver. + /// Convert this component into a channel endpoint and optional owned driver. /// /// The returned [`Channel`] is the canonical frame-aware boundary. It carries /// complete [`TransportFrame`](crate::TransportFrame) values so default @@ -223,7 +204,8 @@ pub trait ConnectTo: Send + 'static { /// /// This method returns: /// - A `Channel` that can be used to communicate with this component - /// - A [`ConnectionDriver`] that drives the component's connection logic + /// - `Some(ConnectionDriver)` when the component owns work to drive + /// - `None` when the channel halves alone own the endpoint's lifetime /// /// 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. @@ -232,17 +214,44 @@ pub trait ConnectTo: Send + 'static { /// /// # Returns /// - /// A tuple of `(Channel, ConnectionDriver)` where the channel is for the caller - /// to use and active drivers must be polled concurrently with channel traffic. - /// Successful active completion ends the endpoint after draining accepted - /// output; passive readiness is not EOF and preserves both channel half-closes. - fn into_channel_and_future(self) -> (Channel, ConnectionDriver) + /// A tuple of `(Channel, Option)`. Owned drivers must be + /// polled concurrently with channel traffic. Successful owned completion + /// ends the endpoint after draining accepted output. `None` is not EOF: + /// preserve both independent channel half-closes. + /// + /// Absence must be handled explicitly; the optional driver is not awaitable: + /// + /// ```compile_fail,E0277 + /// use agent_client_protocol::{Channel, ConnectTo, UntypedRole}; + /// + /// # async fn example() -> agent_client_protocol::Result<()> { + /// let (channel, _peer) = Channel::duplex(); + /// let (_channel, driver) = ConnectTo::::into_channel_and_future(channel); + /// driver.await?; + /// # Ok(()) + /// # } + /// ``` + /// + /// Once present, the owned driver itself is awaitable: + /// + /// ```no_run + /// use agent_client_protocol::{Channel, ConnectionDriver, Result}; + /// + /// async fn drive_owned_work((_channel, driver): (Channel, Option)) -> Result<()> { + /// if let Some(driver) = driver { + /// // In a real adapter, also poll the channel traffic concurrently. + /// driver.await?; + /// } + /// Ok(()) + /// } + /// ``` + fn into_channel_and_future(self) -> (Channel, Option) where Self: Sized, { let (channel_a, channel_b) = Channel::duplex(); let future = ConnectionDriver::new(self.connect_to(channel_b)); - (channel_a, future) + (channel_a, Some(future)) } } @@ -259,7 +268,7 @@ trait ErasedConnectTo: Send { client: Box>, ) -> BoxFuture<'static, Result<()>>; - fn into_channel_and_future_erased(self: Box) -> (Channel, ConnectionDriver); + fn into_channel_and_future_erased(self: Box) -> (Channel, Option); } /// Blanket implementation: any `ConnectTo` can be type-erased. @@ -282,7 +291,7 @@ impl, R: Role> ErasedConnectTo for C { }) } - fn into_channel_and_future_erased(self: Box) -> (Channel, ConnectionDriver) { + fn into_channel_and_future_erased(self: Box) -> (Channel, Option) { (*self).into_channel_and_future() } } @@ -336,7 +345,7 @@ impl ConnectTo for DynConnectTo { .await } - fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { + fn into_channel_and_future(self) -> (Channel, Option) { self.inner.into_channel_and_future_erased() } } @@ -353,29 +362,79 @@ impl Debug for DynConnectTo { mod tests { use super::*; use crate::role::UntypedRole; + use futures::FutureExt as _; + + struct OwnedWork(BoxFuture<'static, Result<()>>); + + impl ConnectTo for OwnedWork { + async fn connect_to(self, _client: impl ConnectTo) -> Result<()> { + self.0.await + } + } #[test] - fn passive_readiness_does_not_claim_an_owned_lifetime() { - let driver = ConnectionDriver::passive(); - assert!(driver.is_passive()); - futures::executor::block_on(driver).unwrap(); + fn raw_channel_has_no_owned_work() { + let (channel, _other) = Channel::duplex(); + let (_, driver) = ConnectTo::::into_channel_and_future(channel); + assert!(driver.is_none()); } #[test] - fn active_driver_preserves_errors_and_polls_unpinned() { + fn owned_driver_preserves_errors_and_polls_unpinned() { let error = crate::Error::internal_error().data("driver failure"); let mut driver = ConnectionDriver::new(futures::future::ready(Err(error.clone()))); - assert!(!driver.is_passive()); assert_eq!(futures::executor::block_on(&mut driver), Err(error)); - // Completion does not reclassify an owned endpoint as passive. - assert!(!driver.is_passive()); + } + + #[test] + fn default_conversion_owns_real_work_until_completion() { + let (done_tx, done_rx) = futures::channel::oneshot::channel(); + let component = OwnedWork(async move { done_rx.await.unwrap() }.boxed()); + let (_channel, driver) = component.into_channel_and_future(); + let mut driver = driver.expect("default conversion always owns its connect_to work"); + assert!((&mut driver).now_or_never().is_none()); + + let error = crate::Error::internal_error().data("owned work failed"); + done_tx.send(Err(error.clone())).unwrap(); + assert_eq!(futures::executor::block_on(driver), Err(error)); + } + + #[test] + fn dropping_optional_owned_driver_cancels_unpolled_work() { + let (done_tx, done_rx) = futures::channel::oneshot::channel::>(); + let component = OwnedWork(async move { done_rx.await.unwrap() }.boxed()); + let (_channel, driver) = component.into_channel_and_future(); + assert!(driver.is_some()); + assert!(!done_tx.is_canceled()); + drop(driver); + assert!(done_tx.is_canceled()); + } + + #[test] + fn type_erasure_preserves_owned_work_and_finish_metadata() { + let outgoing = futures::sink::unfold((), |(), _line: String| { + futures::future::ready(Ok::<_, std::io::Error>(())) + }); + // Independent physical input remains open: only a preserved explicit + // finish handle can complete this driver without read EOF. + let incoming = futures::stream::pending::>(); + let component = DynConnectTo::::new(crate::Lines::new(outgoing, incoming)); + let (_channel, driver) = component.into_channel_and_future(); + let mut driver = driver.expect("erasure must retain ownership"); + assert!((&mut driver).now_or_never().is_none()); + driver + .take_finish() + .expect("erasure must retain physical finish coordination") + .send(()) + .unwrap(); + futures::executor::block_on(driver).unwrap(); } #[test] fn type_erasure_preserves_passive_lifetime() { let (channel, _other) = Channel::duplex(); let (_, driver) = DynConnectTo::::new(channel).into_channel_and_future(); - assert!(driver.is_passive()); + assert!(driver.is_none()); } #[test] diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index dcc175d2..fd455baa 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -1911,8 +1911,7 @@ impl< 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. + // Normalize the transport without losing ownership or finish metadata. let transport_component = crate::DynConnectTo::new(transport); let (transport_channel, transport_future) = transport_component.into_channel_and_future(); let (transport_completion_tx, transport_completion_rx) = oneshot::channel(); @@ -1936,11 +1935,18 @@ impl< pending_replies.registrar(), protocol_mode, ); - let spawn_result = connection.spawn(async move { - let result = transport_future.await; - drop(transport_completion_tx.send(result.clone())); - result - }); + let spawn_result = if let Some(driver) = transport_future { + connection.spawn(async move { + let result = driver.await; + drop(transport_completion_tx.send(result.clone())); + result + }) + } else { + // Channel-only endpoints have no physical sink work to await. + // Their protocol drain marker still orders accepted output. + drop(transport_completion_tx.send(Ok(()))); + Ok(()) + }; // Destructure the channel endpoints let Channel { @@ -6356,8 +6362,9 @@ where } } - fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { - self.into_channel_transport() + fn into_channel_and_future(self) -> (Channel, Option) { + let (channel, driver) = self.into_channel_transport(); + (channel, Some(driver)) } } @@ -6457,7 +6464,7 @@ where ConnectTo::::connect_to(self.into_lines(), client).await } - fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { + fn into_channel_and_future(self) -> (Channel, Option) { ConnectTo::::into_channel_and_future(self.into_lines()) } } @@ -6522,11 +6529,11 @@ impl Channel { /// Passive endpoints instead retain the channel's independent half-close lifetime. pub(crate) async fn copy_with_driver( mut self, - mut driver: crate::ConnectionDriver, + driver: Option, ) -> Result<(), crate::Error> { - if driver.is_passive() { + let Some(mut driver) = driver else { return self.copy().await; - } + }; let mut done = false; loop { @@ -6619,7 +6626,7 @@ 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 passive = client_future.is_none(); let outgoing = Box::pin( Channel { rx: client_channel.rx, @@ -6648,8 +6655,8 @@ impl ConnectTo for Channel { } } - fn into_channel_and_future(self) -> (Channel, crate::ConnectionDriver) { - (self, crate::ConnectionDriver::passive()) + fn into_channel_and_future(self) -> (Channel, Option) { + (self, None) } } diff --git a/src/agent-client-protocol/src/role/acp.rs b/src/agent-client-protocol/src/role/acp.rs index 7892aa24..14e48384 100644 --- a/src/agent-client-protocol/src/role/acp.rs +++ b/src/agent-client-protocol/src/role/acp.rs @@ -970,9 +970,9 @@ async fn reject_initialize( Ok::<_, crate::Error>(()) }; - if future.is_passive() { + let Some(future) = future else { return drain_incoming.await; - } + }; match future::select(future, Box::pin(drain_incoming)).await { future::Either::Left((result, _)) => result, @@ -994,18 +994,24 @@ struct RunningProtocolPeer { enum ProtocolPeerDriver { Passive, Active(crate::ConnectionDriver), - Completed, + Completed { + finish: Option>, + }, } #[cfg(feature = "unstable_protocol_v2")] impl ProtocolPeerDriver { - fn into_driver(self) -> crate::ConnectionDriver { + fn into_driver(self) -> Option { match self { - Self::Passive => crate::ConnectionDriver::passive(), - Self::Active(driver) => driver, + Self::Passive => None, + Self::Active(driver) => Some(driver), // Conversion happens only when handing the peer to its final - // bridge, never while reading its remaining queued frames. - Self::Completed => crate::ConnectionDriver::new(future::ready(Ok(()))), + // bridge, never while reading its remaining queued frames. This + // records actual owned completion, not a passive ready sentinel. + Self::Completed { finish } => Some(match finish { + Some(finish) => crate::ConnectionDriver::with_finish(future::ready(Ok(())), finish), + None => crate::ConnectionDriver::new(future::ready(Ok(()))), + }), } } } @@ -1014,17 +1020,16 @@ impl ProtocolPeerDriver { impl RunningProtocolPeer { fn new(component: impl ConnectTo) -> Self { let (Channel { rx, tx }, future) = component.into_channel_and_future(); - let driver = if future.is_passive() { - ProtocolPeerDriver::Passive - } else { - ProtocolPeerDriver::Active(future) + let driver = match future { + None => ProtocolPeerDriver::Passive, + Some(future) => ProtocolPeerDriver::Active(future), }; Self { rx, tx, driver } } async fn next_frame(self) -> Result, crate::Error> { let Self { mut rx, tx, driver } = self; - let ProtocolPeerDriver::Active(future) = driver else { + let ProtocolPeerDriver::Active(mut future) = driver else { return Ok(rx .next() .await @@ -1033,8 +1038,8 @@ impl RunningProtocolPeer { // Poll the owned driver first: a ready error must not be hidden by // an equally ready frame or clean channel EOF. - match future::select(future, Box::pin(rx.next())).await { - future::Either::Right((Some(frame), future)) => Ok(Some(( + match future::select(&mut future, Box::pin(rx.next())).await { + future::Either::Right((Some(frame), _)) => Ok(Some(( frame, Self { rx, @@ -1042,7 +1047,7 @@ impl RunningProtocolPeer { driver: ProtocolPeerDriver::Active(future), }, ))), - future::Either::Right((None, future)) => { + future::Either::Right((None, _)) => { drop(tx); future.await?; Ok(None) @@ -1061,7 +1066,9 @@ impl RunningProtocolPeer { Self { rx, tx, - driver: ProtocolPeerDriver::Completed, + driver: ProtocolPeerDriver::Completed { + finish: future.take_finish(), + }, }, ))) } @@ -1147,10 +1154,14 @@ async fn pipe_protocol_peers_until_done( ) -> Result<(), crate::Error> { let mut left_driver = left.driver.into_driver(); let mut right_driver = right.driver.into_driver(); - let left_passive = left_driver.is_passive(); - let right_passive = right_driver.is_passive(); - let left_finish = left_driver.take_finish(); - let right_finish = right_driver.take_finish(); + let left_passive = left_driver.is_none(); + let right_passive = right_driver.is_none(); + let left_finish = left_driver + .as_mut() + .and_then(crate::ConnectionDriver::take_finish); + let right_finish = right_driver + .as_mut() + .and_then(crate::ConnectionDriver::take_finish); let left_to_right = Box::pin( Channel { rx: left.rx, @@ -1822,6 +1833,43 @@ mod lifetime_tests { assert!(matches!(peer.driver, ProtocolPeerDriver::Passive)); } + #[tokio::test] + async fn owned_peer_preserves_finish_metadata_through_active_and_completed_states() { + let (Channel { rx, tx }, remote) = Channel::duplex(); + let (done_tx, done_rx) = futures::channel::oneshot::channel(); + let (finish_tx, finish_rx) = futures::channel::oneshot::channel(); + let peer = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(ConnectionDriver::with_finish( + async move { + done_rx.await.unwrap(); + Ok(()) + }, + finish_tx, + )), + }; + + remote.tx.unbounded_send(frame()).unwrap(); + let (_, peer) = peer.next_frame().await.unwrap().unwrap(); + assert!(matches!(&peer.driver, ProtocolPeerDriver::Active(_))); + + remote.tx.unbounded_send(frame()).unwrap(); + done_tx.send(()).unwrap(); + let (_, peer) = peer.next_frame().await.unwrap().unwrap(); + assert!(matches!( + &peer.driver, + ProtocolPeerDriver::Completed { finish: Some(_) } + )); + let mut driver = peer + .driver + .into_driver() + .expect("completed owned work must not become passive"); + driver.take_finish().unwrap().send(()).unwrap(); + finish_rx.await.unwrap(); + driver.await.unwrap(); + } + #[tokio::test] async fn active_initialization_drains_accepted_frames_without_escaped_sender_eof() { let (Channel { rx, tx }, remote) = Channel::duplex(); @@ -1895,7 +1943,10 @@ mod lifetime_tests { } Ok(()) }; - crate::util::run_until(driver, foreground).await + match driver { + Some(driver) => crate::util::run_until(driver, foreground).await, + None => foreground.await, + } } } diff --git a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs index e56f778e..aca89329 100644 --- a/src/agent-client-protocol/tests/jsonrpc_transport_close.rs +++ b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs @@ -108,7 +108,10 @@ impl ConnectTo for QueuedClient { drop(self.escaped.send(channel.tx.clone())); let _ = self.started.send(()); drop(channel); - transport_future.await + if let Some(driver) = transport_future { + driver.await?; + } + Ok(()) } } @@ -184,7 +187,9 @@ impl ConnectTo for CompletingClient { .map_err(Error::into_internal_error)?; } drop(channel); - driver.await?; + if let Some(driver) = driver { + driver.await?; + } self.0 } } @@ -215,7 +220,10 @@ impl ConnectTo for RequestReplyClient { assert_eq!(id, RequestId::Number(42)); assert_eq!(result, serde_json::json!({ "status": "received" })); drop(channel); - driver.await + if let Some(driver) = driver { + driver.await?; + } + Ok(()) } } @@ -233,8 +241,8 @@ impl ConnectTo for DrivenEndpoint { Ok(()) } - fn into_channel_and_future(self) -> (Channel, ConnectionDriver) { - (self.channel, self.driver) + fn into_channel_and_future(self) -> (Channel, Option) { + (self.channel, Some(self.driver)) } } @@ -430,7 +438,7 @@ async fn assert_passive_bridge_preserves_half_close( } #[test] -fn channel_and_erased_channel_drivers_are_explicitly_passive() { +fn channel_and_erased_channel_have_no_driver() { for erased in [false, true] { let (endpoint, _peer) = Channel::duplex(); let (_channel, driver) = if erased { @@ -438,18 +446,13 @@ fn channel_and_erased_channel_drivers_are_explicitly_passive() { } else { ConnectTo::::into_channel_and_future(endpoint) }; - assert!(driver.is_passive()); - driver - .now_or_never() - .expect("passive driver should be immediately ready") - .expect("passive readiness should succeed"); + assert!(driver.is_none(), "Channel does not own runnable work"); } +} +#[test] +fn ready_owned_driver_is_awaitable() { let driver = ConnectionDriver::new(future::ready(Ok(()))); - assert!( - !driver.is_passive(), - "ready owned work is active even when it completes immediately" - ); driver.now_or_never().unwrap().unwrap(); } @@ -679,6 +682,7 @@ async fn transport_channel_keeps_read_half_open_after_write_half_closes() { let (mut peer_outgoing, sdk_incoming) = tokio::io::duplex(1024); let transport = ByteStreams::new(sdk_outgoing.compat_write(), sdk_incoming.compat()); let (channel, transport_future) = ConnectTo::::into_channel_and_future(transport); + let transport_future = transport_future.expect("byte streams own a transport driver"); let Channel { mut rx, tx } = channel; tx.unbounded_send(TransportFrame::Single( diff --git a/src/agent-client-protocol/tests/protocol_v2.rs b/src/agent-client-protocol/tests/protocol_v2.rs index 8c316339..6cba4bf9 100644 --- a/src/agent-client-protocol/tests/protocol_v2.rs +++ b/src/agent-client-protocol/tests/protocol_v2.rs @@ -405,7 +405,7 @@ impl ConnectTo for InitializingV2Client { impl ConnectTo for FutureInitializeV2Client { async fn connect_to(self, agent: impl ConnectTo) -> Result<(), Error> { let (mut channel, agent_future) = ConnectTo::::into_channel_and_future(agent); - let agent_task = tokio::spawn(agent_future); + let agent_task = agent_future.map(tokio::spawn); channel .tx @@ -426,11 +426,15 @@ impl ConnectTo for FutureInitializeV2Client { }; let initialize = v2::InitializeResponse::from_value("initialize", result)?; assert_eq!(initialize.protocol_version, ProtocolVersion::V2); - agent_task.abort(); + if let Some(task) = agent_task { + task.abort(); + } return Ok(()); } - agent_task.abort(); + if let Some(task) = agent_task { + task.abort(); + } Err(agent_client_protocol::util::internal_error( "v2 agent did not respond to initialize", )) @@ -485,7 +489,7 @@ async fn assert_malformed_initialize_rejected(params: Map) -> Res agent_client_protocol::on_receive_request!(), ); let (mut channel, agent_future) = ConnectTo::::into_channel_and_future(agent); - let agent_task = tokio::spawn(agent_future); + let agent_task = tokio::spawn(agent_future.expect("v2 agent owns a connection driver")); channel .tx @@ -2024,7 +2028,7 @@ async fn protocol_router_v2_only_rejects_v1_client() -> Result<(), Error> { )); let (mut channel, agent_future) = ConnectTo::::into_channel_and_future(agent); - let agent_task = tokio::spawn(agent_future); + let agent_task = tokio::spawn(agent_future.expect("agent router owns a connection driver")); channel .tx @@ -2760,7 +2764,7 @@ async fn protocol_router_routes_future_protocol_version_to_v2() -> Result<(), Er )); let (mut channel, agent_future) = ConnectTo::::into_channel_and_future(agent); - let agent_task = tokio::spawn(agent_future); + let agent_task = tokio::spawn(agent_future.expect("agent router owns a connection driver")); let mut initialize = json_value(v2_initialize_request(ProtocolVersion::from(3_u16)))?; initialize 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 ffa4e5cb..bc43ce39 100644 --- a/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs +++ b/src/agent-client-protocol/tests/proxy_protocol_router_v2.rs @@ -75,7 +75,7 @@ async fn request( params: Value, ) -> Result, Error> { let (Channel { mut rx, tx }, future) = ConnectTo::::into_channel_and_future(router); - let task = tokio::spawn(future); + let task = tokio::spawn(future.expect("proxy router owns a connection driver")); let request_id = v1::RequestId::Number(1); tx.unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( From ad9c495ed22e71b46b8575713ab020ff25ebe886 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Fri, 2 Oct 2026 16:50:39 +0200 Subject: [PATCH 3/3] feat(acp): support graceful driver completion and composition Expose cooperative finish requests, preserve their contract after an earlier request or future decoration, and give wrappers a map_future path. Apply one drain rule to builder and protocol-router entry points, keep physical write-half shutdown and late-input error handling, and retain HTTP router cancellation ownership. Add public custom-adapter, direct/normalized, and gated-drain regressions and migration guidance. --- md/migration-connection-drivers.md | 153 ++- md/transport-architecture.md | 23 +- src/agent-client-protocol-http/CHANGELOG.md | 3 + .../src/connection.rs | 105 ++- src/agent-client-protocol/CHANGELOG.md | 28 + src/agent-client-protocol/src/component.rs | 252 ++++- src/agent-client-protocol/src/jsonrpc.rs | 395 ++++++-- .../src/jsonrpc/incoming_actor.rs | 27 +- .../src/jsonrpc/outgoing_actor.rs | 41 +- .../src/jsonrpc/transport_actor.rs | 126 ++- src/agent-client-protocol/src/role/acp.rs | 327 ++++++- .../tests/connection_driver_normalization.rs | 720 +++++++++++++++ .../tests/cooperative_connection_driver.rs | 871 ++++++++++++++++++ .../tests/protocol_driver_finish.rs | 383 ++++++++ 14 files changed, 3287 insertions(+), 167 deletions(-) create mode 100644 src/agent-client-protocol/tests/connection_driver_normalization.rs create mode 100644 src/agent-client-protocol/tests/cooperative_connection_driver.rs create mode 100644 src/agent-client-protocol/tests/protocol_driver_finish.rs diff --git a/md/migration-connection-drivers.md b/md/migration-connection-drivers.md index 23042c2a..86aaf39c 100644 --- a/md/migration-connection-drivers.md +++ b/md/migration-connection-drivers.md @@ -13,7 +13,10 @@ introduces no new frame-size, queue, or task limits. If your component implements only `connect_to`, no change is needed. The default conversion still creates a channel pair and drives your component, now -returning `Some(ConnectionDriver)`. +returning `Some(ConnectionDriver)`. That default wraps opaque work; it cannot +infer a physical finish hook. A buffered transport that needs a finite +foreground to await physical flush should override normalization with +`with_finish`, as described below. Low-level callers must handle the optional work explicitly. The optional value is not a future: awaiting it directly no longer compiles. For a component that @@ -54,29 +57,165 @@ fn into_channel_and_future(self) -> (Channel, Option) { An existing `Channel` has no driver. There is no awaitable passive sentinel, and no finish hook belongs to the `None` case. `ConnectionDriver` always holds -real owned work; some built-in owned drivers additionally support private -finish coordination. Passive bridges retain each read/write half until its +real owned work; cooperative drivers additionally support a finish hook. +Passive bridges retain each read/write half until its own closure; an input half-close can still be followed by a final response. If a wrapper simply exposes another component's endpoint, return its original `(channel, optional_driver)` pair. Re-boxing an owned driver and wrapping it -with `new` would erase its built-in finish coordination. Inventing a ready +with `new` would hide its finish capability. Inventing a ready driver for `None` would also turn absence into a false completion signal. +For tracing, error annotation, or completion cleanup, decorate the future with +`map_future`. This preserves both the finish capability and any already-issued +request; opaque work stays opaque: + +```rust,ignore +use futures::FutureExt; + +let (channel, driver) = component.into_channel_and_future(); +let driver = driver.map(|driver| { + driver.map_future(|work| { + work.inspect(|result| eprintln!("transport completed: {result:?}")) + }) +}); +(channel, driver) +``` + +The transformed future must still drive the original work and must not report +success before its accepted output has drained. + ## Completion and drain responsibilities Poll owned work and outbound forwarding concurrently. A driver may need its outbound request to be delivered before it can receive a response and finish. An adapter must not report success before flushing output it already accepted. -Custom normalized drivers should finish their own work and drain output after -their channel input closes. The SDK cannot infer how to flush an arbitrary -opaque future or external buffer. +Use `ConnectionDriver::with_finish(future, finish)` for a custom normalized +transport that needs to flush during finite foreground shutdown. The +nonblocking `FnOnce()` hook requests graceful completion; the future proves +completion and reports any I/O error. + +```rust,ignore +let (finish_tx, finish_rx) = futures::channel::oneshot::channel(); +let future = async move { + // Keep processing input and output while waiting for the finish request. + // A dropped sender is not a finish request; it may simply mean that + // finish control was abandoned while normal half-closes remain in use. + // + // After a successful signal, seal the outgoing queue, drain every accepted + // frame, and flush/close the physical write half. Do not wait for remote + // read EOF; continue observing genuine I/O errors during the drain. + run_custom_adapter(outgoing_rx, physical_io, finish_rx).await +}; +let driver = ConnectionDriver::with_finish(future, move || { + let _ = finish_tx.send(()); +}); +(channel, Some(driver)) +``` + +SDK shutdown coordination invokes this hook only after protocol output has +been handed off to the normalized transport, then awaits the driver. Low-level +callers can use `driver.request_finish()` themselves. A `true` return means +cooperative finish is supported and has been requested, including a request +already issued. Requests are idempotent, but the hook runs only once. A `false` +return means opaque work, not "already requested." + +Requesting finish does not prove output has finished flushing; continue polling +or await the driver. Capability remains intact if that driver is handed to +another owner while flushing. Quiesce and hand off output before requesting +finish; idempotence does not permit new output after sealing. Dropping the +driver drops its owned future without a graceful request. Dropping only the +hook does not invoke it or necessarily stop that future. + +There is no implicit timeout. A cooperative adapter that cannot flush keeps +the connection pending, so applications that need a deadline must impose one +and accept that cancelling it can truncate output. `with_finish` declares the +adapter's contract; it cannot make an arbitrary future or external buffer +flush automatically. Built-in `Lines` and `ByteStreams` preserve normal half-close behavior. When their owner explicitly finishes, they drain accepted output while continuing to poll incoming I/O for errors, rather than waiting for unrelated remote input to reach EOF. Errors may terminate the connection without graceful drain. +## Direct adapter entry point + +Returning a cooperative driver from `into_channel_and_future` lets normalized +SDK consumers coordinate finish. A custom transport's direct `connect_to` +implementation must coordinate it too: `try_join!(bridge, driver)` alone can +wait forever after a finite peer has returned. + +This scaffold follows the built-in `Lines` policy, using only public APIs: + +```rust,ignore +use agent_client_protocol::{Channel, ConnectTo, ConnectionDriver, Result, UntypedRole}; +use futures::{future::{select, Either}, FutureExt}; + +struct BufferedAdapter { + channel: Channel, + driver: ConnectionDriver, +} + +impl ConnectTo for BufferedAdapter { + async fn connect_to(self, peer: impl ConnectTo) -> Result<()> { + let bridge = Box::pin(self.channel.connect_to(peer)); + match select(bridge, self.driver).await { + Either::Left((result, mut driver)) => { + result?; // The peer's accepted output has been handed off. + if driver.request_finish() { + driver.await // Prove physical drain; propagate its errors. + } else { + // Preserve a ready error, then cancel opaque work. + driver.now_or_never().unwrap_or(Ok(())) + } + } + Either::Right((result, _bridge)) => result, + } + } + + fn into_channel_and_future(self) -> (Channel, Option) { + (self.channel, Some(self.driver)) + } +} +``` + +The adapter's own future must own/close its physical producers before reporting +completion. Its finish implementation must stop forwarding successful input +to a completed peer while still observing genuine read errors during drain. +Otherwise, late input can fail against the dropped receiver and cancel final +output. Both direct and normalized entry points should be tested with output +backpressure and independently open input. + +## Finite foreground shutdown + +On successful `Builder::connect_with` foreground completion, routable queued +output is drained. Requests still blocked on unresolved readiness are failed +and removed rather than published after shutdown. Physical transport and +protocol progress continue during this drain; queued application tasks are +not started merely to finish the sink. The inherited cleanup coordinator can +still poll application tasks while protecting a close callback already +underway; a blocked close callback can delay completion. + +Foreground success stops beginning new application delivery or close +callbacks. Physical reads remain driven without delivering their input to the +completed foreground. A close callback already underway finishes before the +outgoing drain boundary seals, and its errors retain precedence. + +Protocol connectors and routers use the same ownership-aware rule. An owned +foreground's completion requests cooperative drain instead of waiting for +unrelated remote input. Initialization rejection also hands off its reply +before requesting finish. Passive half-closes alone do not request finish; +they preserve the other direction for a final response. + +Cooperative drivers, both built-in and custom, are awaited through physical +write shutdown. +An opaque driver constructed with `ConnectionDriver::new(future)` has no +externally requestable finish control: finite foreground shutdown transfers +protocol output into its normalized channel, then cancels that work without +guaranteeing custom physical flush. Reactive `connect_to` still joins owned +work after input EOF. Choose `new` for opaque cancellable work and `with_finish` +when the adapter can honor an explicit graceful-finish request. + See [Transport Architecture](./transport-architecture.md#component-boundary) for the active/passive boundary and forwarding rules. diff --git a/md/transport-architecture.md b/md/transport-architecture.md index 1013846d..994d58dc 100644 --- a/md/transport-architecture.md +++ b/md/transport-architecture.md @@ -292,15 +292,30 @@ a component waiting for a response to its own request could deadlock. Buffered adapters are responsible for flushing their accepted output before reporting successful completion. The built-in line and byte-stream adapters keep the read half moving during write drain and propagate incoming errors; -their explicit finish handling does not require remote read EOF. Merely -wrapping an arbitrary future cannot make an opaque custom adapter drain safely. +their explicit finish handling does not require remote read EOF. +`ConnectionDriver::with_finish(future, finish)` lets custom adapters declare +the same cooperative contract. Its one-shot, nonblocking hook requests finish +after protocol output handoff; the still-polled future performs the drain and +reports completion or I/O errors. `request_finish()` exposes this request to +low-level callers and is idempotent: a supported request remains supported +after it has been issued. The callback still runs at most once. Neither +requesting finish nor dropping the hook proves a successful flush, and no +implicit timeout is imposed. + +`ConnectionDriver::new(future)` remains appropriate for opaque cancellable +work. A finite foreground does not wait indefinitely for such work after +handing off protocol output. Merely wrapping an arbitrary future cannot make +an opaque custom adapter drain safely. Most components implement only `connect_to`; default normalization supplies `Some(driver)` containing the owned work. Direct transports override `into_channel_and_future` to avoid an intermediate copy. Wrappers that expose an existing endpoint should forward -its driver unchanged so passive identity and built-in completion handling are -not lost. +its driver unchanged so absence and cooperative completion handling are +not lost. Wrappers that decorate execution use `map_future` to transform the +owned future while preserving its finish capability and requested state. +Direct custom transport entry points must also request and await cooperative +finish after a finite peer completes; a bare join does not provide that step. See [Migrating Connection Drivers](./migration-connection-drivers.md) for custom override changes. This lifecycle distinction does not change the existing raw diff --git a/src/agent-client-protocol-http/CHANGELOG.md b/src/agent-client-protocol-http/CHANGELOG.md index c52700bd..8d2efb1e 100644 --- a/src/agent-client-protocol-http/CHANGELOG.md +++ b/src/agent-client-protocol-http/CHANGELOG.md @@ -18,6 +18,9 @@ - On active agent completion, reject further output from escaped sender clones and drain accepted frames before removing the connection and closing its streams, without waiting for those clones to be dropped. +- Keep router cancellation owned while natural cleanup awaits its drain. + Explicit shutdown no longer detaches a taken router task or retains the + connection through that orphaned task. ## [2.2.0](https://github.com/agentclientprotocol/rust-sdk/compare/agent-client-protocol-http-v2.1.0...agent-client-protocol-http-v2.2.0) - 2026-09-18 diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index 8e190c8c..04df3888 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -589,13 +589,24 @@ impl ConnectionRegistry { } } +// Taking a JoinHandle transfers cancellation ownership out of the connection. +// Keep that ownership through the await: dropping a bare handle detaches it. +struct AbortTakenRouterOnDrop(tokio::task::AbortHandle); + +impl Drop for AbortTakenRouterOnDrop { + fn drop(&mut self) { + self.0.abort(); + } +} + async fn drain_connection_router(connection: Weak) -> Option> { let connection = connection.upgrade()?; let router_handle = connection.router_handle.lock().await.take(); - if let Some(h) = router_handle - && let Err(error) = h.await - { - error!("outbound router task failed while draining: {error}"); + if let Some(handle) = router_handle { + let _abort_on_drop = AbortTakenRouterOnDrop(handle.abort_handle()); + if let Err(error) = handle.await { + error!("outbound router task failed while draining: {error}"); + } } Some(connection) } @@ -1151,6 +1162,92 @@ mod tests { )); } + struct RouterDropProbe(Option>); + + impl Drop for RouterDropProbe { + fn drop(&mut self) { + let _sent = self.0.take().unwrap().send(()); + } + } + + #[tokio::test] + async fn shutdown_during_natural_router_drain_cancels_owned_router() { + let emit = Arc::new(Notify::new()); + let registry = ConnectionRegistry::new(Arc::new(FinalFrameThenExitAgentFactory { + emit: emit.clone(), + escaped_output: Arc::new(StdMutex::new(None)), + })); + let (connection_id, connection) = registry.create_connection().await; + let weak_connection = Arc::downgrade(&connection); + let mut outbound = connection.subscribe_connection_stream().unwrap(); + let mut closed = connection.subscribe_closed(); + let mut frames = connection.outbound_rx.lock().await.take().unwrap(); + let (routing_started_tx, routing_started_rx) = tokio::sync::oneshot::channel(); + let (release_router_tx, release_router_rx) = tokio::sync::oneshot::channel::<()>(); + let (dropped_tx, mut dropped_rx) = tokio::sync::oneshot::channel(); + let routing_connection = connection.clone(); + let router = tokio::spawn(async move { + let _drop_probe = RouterDropProbe(Some(dropped_tx)); + let frame = frames.recv().await.expect("accepted final frame"); + routing_started_tx.send(()).unwrap(); + let _released = release_router_rx.await; + routing_connection.route_outbound(frame).await.unwrap(); + while let Some(frame) = frames.recv().await { + routing_connection.route_outbound(frame).await.unwrap(); + } + }); + // Keep an independent abort handle solely to clean up a failing probe. + let failed_test_cleanup = router.abort_handle(); + *connection.router_handle.lock().await = Some(router); + + emit.notify_one(); + timeout(Duration::from_secs(1), async { + routing_started_rx.await.unwrap(); + while connection.router_handle.lock().await.is_some() { + tokio::task::yield_now().await; + } + }) + .await + .expect("natural cleanup must have taken the router join handle"); + assert!(registry.get(&connection_id).await.is_some()); + assert!(!*closed.borrow()); + + // Match DELETE: remove the discoverable connection and shut it down. + let removed = registry.remove(&connection_id).await.unwrap(); + removed.shutdown().await; + closed.changed().await.unwrap(); + assert!(*closed.borrow()); + // Raw mailbox EOF is not the stream-closure contract. + assert_eq!(outbound.try_recv(), Err(mpsc::error::TryRecvError::Empty)); + drop(removed); + drop(connection); + + // The gate remains held throughout both observations. Only cancellation, + // not releasing or dropping its sender, may finish the owned router. + let router_cancelled = timeout(Duration::from_secs(1), &mut dropped_rx) + .await + .is_ok_and(|result| result.is_ok()); + let connection_released = timeout(Duration::from_secs(1), async { + while weak_connection.upgrade().is_some() { + tokio::task::yield_now().await; + } + }) + .await + .is_ok(); + if !router_cancelled { + failed_test_cleanup.abort(); + timeout(Duration::from_secs(1), &mut dropped_rx) + .await + .expect("failed-test cleanup must cancel the orphan") + .unwrap(); + } + drop(release_router_tx); + assert!( + router_cancelled && connection_released, + "shutdown must cancel the gated owned router without releasing its gate: router_cancelled={router_cancelled}, connection_released={connection_released}" + ); + } + #[tokio::test] async fn protocol_level_notification_routes_to_connection_stream() { let exit = Arc::new(Notify::new()); diff --git a/src/agent-client-protocol/CHANGELOG.md b/src/agent-client-protocol/CHANGELOG.md index cc659f1c..9a7cb050 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -20,6 +20,13 @@ - Add an owned-only, awaitable `ConnectionDriver`. Passive endpoints have no driver, so absence cannot be mistaken for a successful completed future. +- Add `ConnectionDriver::with_finish(future, finish)` and `request_finish()` for + custom adapters to request graceful shutdown separately from awaiting + completion. Requests are idempotent; the one-shot hook signals the adapter + once, and its future must drain, flush, and close accepted output before + reporting success. Already-requested drivers retain that contract on handoff. +- Add `ConnectionDriver::map_future` for tracing, error annotation, and + completion cleanup without losing finish capability or requested state. - Add a default-enabled `schemars` feature that forwards JSON Schema support to the schema crate and gates the typed MCP tool helpers. Set `default-features = false` to use the core SDK without `schemars`; custom MCP @@ -34,6 +41,27 @@ - Drain accepted output when an owned component finishes, including through line and byte-stream adapters, without requiring unrelated remote input to close. Ready component and I/O failures remain authoritative during drain. +- Drain routable queued output on successful `Builder::connect_with` + foreground completion and finish cooperative physical transports without + starting queued application tasks solely for physical drain. Preserve the + inherited underway-close-callback cleanup phase. Fail unresolved + request-readiness gates rather than waiting indefinitely or publishing those + requests after shutdown. +- Close and drain the incoming producer boundary when owned transport work + completes, including when escaped senders remain alive. +- Forward physical write-half shutdown through byte-stream adapters so split + streams can receive a final reverse response after output EOF. Preserve + partial-write, flush, pending-shutdown, and close-error behavior. +- Discard input to a completed foreground during physical finish, including + frames already queued at normalization, without losing genuine I/O errors + or cancelling accepted output. +- Apply one ownership-aware drain rule to protocol connectors and agent/proxy + routers. Request cooperative finish after initialization rejection or owned + foreground completion; do not indefinitely join opposed opaque work. +- Do not begin late application delivery or close callbacks after finite + foreground success. An underway close callback completes before output is + sealed, so clean EOF during physical drain cannot turn a late callback send + into a closed-queue failure that cancels accepted output. ## [2.2.0](https://github.com/agentclientprotocol/rust-sdk/compare/v2.1.0...v2.2.0) - 2026-09-18 diff --git a/src/agent-client-protocol/src/component.rs b/src/agent-client-protocol/src/component.rs index 72373f13..c5bcb6bf 100644 --- a/src/agent-client-protocol/src/component.rs +++ b/src/agent-client-protocol/src/component.rs @@ -37,27 +37,56 @@ use std::{ use crate::{Channel, Result, role::Role}; +// Presence of the control records the cooperative contract, even after its +// one-shot action has run. Moving a requested driver must not erase that fact. +pub(crate) struct FinishControl { + hook: Option>, +} + +impl FinishControl { + fn new(hook: impl FnOnce() + Send + 'static) -> Self { + Self { + hook: Some(Box::new(hook)), + } + } + + pub(crate) fn request(&mut self) { + if let Some(hook) = self.hook.take() { + hook(); + } + } +} + /// Drives owned endpoint work. /// /// A driver owns the endpoint: successful completion means no further /// output is expected, and adapters must drain output already accepted before -/// terminating. An error terminates the connection immediately. +/// terminating. Errors may abort the connection without guaranteed output +/// drain; an adapter may still preserve queued error replies before terminating. /// /// Poll the driver concurrently with channel traffic. Endpoints without owned /// work return `None` from [`ConnectTo::into_channel_and_future`], not a driver: /// their channel halves independently determine their lifetime. +/// +/// Use [`new`](Self::new) for opaque work or +/// [`with_finish`](Self::with_finish) for a transport that can finish gracefully. +/// Use [`map_future`](Self::map_future) to decorate existing work without losing +/// its finish capability. #[must_use = "connection drivers must be polled to make progress"] pub struct ConnectionDriver { future: BoxFuture<'static, Result<()>>, - finish: Option>, + finish: Option, } impl ConnectionDriver { /// Create a driver that owns the endpoint's lifetime. /// - /// Custom normalized adapters must finish accepted output when their channel - /// input closes. This constructor cannot externally flush or shut down an - /// opaque future that waits for additional, independently owned input. + /// This driver has no cooperative finish hook. A finite foreground may drop + /// it after handing off accepted output, rather than wait for arbitrary + /// work to finish. Reactive serving still awaits owned work after input EOF. + /// + /// Custom transports that need to flush before a finite foreground returns + /// should use [`with_finish`](Self::with_finish) instead. pub fn new(future: impl Future> + Send + 'static) -> Self { Self { future: Box::pin(future), @@ -65,19 +94,115 @@ impl ConnectionDriver { } } - // Physical transports can finish their write half without waiting for read - // EOF. Keep this coordination private; arbitrary futures cannot support it. - pub(crate) fn with_finish( + /// Create owned work that supports cooperative graceful completion. + /// + /// The finish hook only requests completion; it must be nonblocking and + /// should signal the future to stop accepting output, drain what it has + /// already accepted, flush and close its write half, then return. It must + /// not require independently open remote input to reach EOF. The future + /// remains responsible for reporting I/O and flush errors. + /// + /// SDK consumers invoke the hook after handing off their accepted output, + /// then continue polling the driver until completion. There is no implicit + /// timeout: if the adapter cannot finish, the enclosing connection remains + /// pending and may be cancelled by its caller. + /// + /// The hook is invoked at most once. Dropping the driver drops its owned + /// future without requesting graceful completion. Dropping only the hook + /// does not invoke it or necessarily stop the work. + /// + /// # Example + /// + /// A custom adapter can use any signal understood by its future. For + /// example, a one-shot channel separates the finish request from completion: + /// + /// ``` + /// use agent_client_protocol::ConnectionDriver; + /// use futures::{channel::oneshot, FutureExt}; + /// + /// let (finish_tx, finish_rx) = oneshot::channel(); + /// let mut driver = ConnectionDriver::with_finish( + /// async move { + /// if finish_rx.await.is_err() { + /// // Losing the hook must not masquerade as a finish request. + /// futures::future::pending::<()>().await; + /// } + /// // Seal the adapter's outgoing queue, drain it, and flush/close + /// // the physical writer here before returning. + /// Ok(()) + /// }, + /// move || { let _ = finish_tx.send(()); }, + /// ); + /// + /// assert!((&mut driver).now_or_never().is_none()); + /// assert!(driver.request_finish()); + /// assert!(driver.request_finish()); // Supported, but the hook runs only once. + /// futures::executor::block_on(driver).unwrap(); + /// ``` + pub fn with_finish( future: impl Future> + Send + 'static, - finish: futures::channel::oneshot::Sender<()>, + finish: impl FnOnce() + Send + 'static, ) -> Self { Self { future: Box::pin(future), - finish: Some(finish), + finish: Some(FinishControl::new(finish)), } } - pub(crate) fn take_finish(&mut self) -> Option> { + /// Decorate the owned future while preserving its finish capability. + /// + /// This is useful for tracing, error annotation, or completion cleanup. + /// Wrapping this driver in [`new`](Self::new) instead would hide its finish + /// control from the outer driver. + /// + /// `map` is called immediately and receives the boxed future, not the + /// driver. Its returned future must uphold the same completion contract: + /// keep driving the original work and do not report success before accepted + /// output is drained. An already-requested finish remains requested, and + /// opaque work remains opaque. + /// + /// ``` + /// use agent_client_protocol::ConnectionDriver; + /// use futures::FutureExt; + /// + /// let driver = ConnectionDriver::new(async { Ok(()) }); + /// let decorated = driver.map_future(|work| { + /// work.inspect(|result| eprintln!("transport completed: {result:?}")) + /// }); + /// futures::executor::block_on(decorated).unwrap(); + /// ``` + pub fn map_future(self, map: impl FnOnce(BoxFuture<'static, Result<()>>) -> F) -> Self + where + F: Future> + Send + 'static, + { + Self { + future: Box::pin(map(self.future)), + finish: self.finish, + } + } + + /// Request graceful completion, without waiting for it. + /// + /// Returns `true` if this driver supports cooperative finish, including + /// when finish was already requested. Repeated requests are idempotent: + /// the hook runs at most once and the driver retains its graceful-finish + /// contract across wrapping or ownership handoff. + /// + /// Returns `false` for opaque work constructed with [`new`](Self::new); + /// this method does not cancel that work. A `true` return does not prove + /// flushing is complete: continue polling or await the driver to observe + /// completion and any errors. + #[must_use] + pub fn request_finish(&mut self) -> bool { + if let Some(finish) = self.finish.as_mut() { + finish.request(); + true + } else { + false + } + } + + pub(crate) fn take_finish(&mut self) -> Option { self.finish.take() } } @@ -386,6 +511,102 @@ mod tests { assert_eq!(futures::executor::block_on(&mut driver), Err(error)); } + #[test] + fn finish_request_is_idempotent_and_does_not_mean_completion() { + let (finish_tx, finish_rx) = futures::channel::oneshot::channel(); + let (flushed_tx, flushed_rx) = futures::channel::oneshot::channel(); + let mut driver = ConnectionDriver::with_finish( + async move { + finish_rx.await.unwrap(); + flushed_rx.await.unwrap() + }, + move || finish_tx.send(()).unwrap(), + ); + + assert!((&mut driver).now_or_never().is_none()); + assert!(driver.request_finish()); + assert!(driver.request_finish()); + assert!((&mut driver).now_or_never().is_none()); + + let error = crate::Error::internal_error().data("custom flush failed"); + flushed_tx.send(Err(error.clone())).unwrap(); + assert_eq!(futures::executor::block_on(driver), Err(error)); + } + + #[test] + fn opaque_driver_cannot_be_cooperatively_finished() { + let mut driver = ConnectionDriver::new(futures::future::pending()); + assert!(!driver.request_finish()); + assert!((&mut driver).now_or_never().is_none()); + } + + #[test] + fn future_decoration_preserves_finish_and_completion_errors() { + use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }; + + for request_before_wrapping in [false, true] { + let (finish_tx, finish_rx) = futures::channel::oneshot::channel(); + let (flush_tx, flush_rx) = futures::channel::oneshot::channel(); + let calls = Arc::new(AtomicUsize::new(0)); + let hook_calls = calls.clone(); + let mut driver = ConnectionDriver::with_finish( + async move { + finish_rx.await.unwrap(); + flush_rx.await.unwrap() + }, + move || { + hook_calls.fetch_add(1, Ordering::SeqCst); + finish_tx.send(()).unwrap(); + }, + ); + if request_before_wrapping { + assert!(driver.request_finish()); + } + + let observed = Arc::new(AtomicUsize::new(0)); + let observe_completion = observed.clone(); + let mut decorated = driver.map_future(|work| { + work.inspect(move |_| { + observe_completion.fetch_add(1, Ordering::SeqCst); + }) + }); + assert!(decorated.request_finish()); + assert!(decorated.request_finish()); + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert!((&mut decorated).now_or_never().is_none()); + assert_eq!(observed.load(Ordering::SeqCst), 0); + + let error = crate::Error::internal_error().data("decorated flush failed"); + flush_tx.send(Err(error.clone())).unwrap(); + assert_eq!(futures::executor::block_on(decorated), Err(error)); + assert_eq!(observed.load(Ordering::SeqCst), 1); + } + } + + #[test] + fn future_decoration_does_not_make_opaque_work_cooperative() { + let driver = ConnectionDriver::new(futures::future::pending()); + let mut decorated = driver.map_future(|work| work); + + assert!(!decorated.request_finish()); + assert!((&mut decorated).now_or_never().is_none()); + } + + #[test] + fn dropping_driver_does_not_invoke_finish_hook() { + let invoked = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); + let hook_invoked = invoked.clone(); + let driver = ConnectionDriver::with_finish(futures::future::pending(), move || { + hook_invoked.store(true, std::sync::atomic::Ordering::Release); + }); + + drop(driver); + assert!(!invoked.load(std::sync::atomic::Ordering::Acquire)); + } + #[test] fn default_conversion_owns_real_work_until_completion() { let (done_tx, done_rx) = futures::channel::oneshot::channel(); @@ -422,11 +643,10 @@ mod tests { let (_channel, driver) = component.into_channel_and_future(); let mut driver = driver.expect("erasure must retain ownership"); assert!((&mut driver).now_or_never().is_none()); - driver - .take_finish() - .expect("erasure must retain physical finish coordination") - .send(()) - .unwrap(); + assert!( + driver.request_finish(), + "erasure must retain finish coordination" + ); futures::executor::block_on(driver).unwrap(); } diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index fd455baa..4bf12930 100644 --- a/src/agent-client-protocol/src/jsonrpc.rs +++ b/src/agent-client-protocol/src/jsonrpc.rs @@ -1804,9 +1804,9 @@ impl< self, transport: impl ConnectTo + 'static, ) -> Result<(), crate::Error> { - let (_, future) = self.into_connection_and_future(transport, async move |cx| { + let (_, future) = self.into_connection_and_future(transport, true, async move |cx| { cx.incoming_closed().await; - cx.drain_outgoing().await + Ok(()) }); future.await } @@ -1881,9 +1881,10 @@ impl< transport: impl ConnectTo + 'static, main_fn: impl AsyncFnOnce(Context::Connection) -> Result, ) -> Result { - let (_, future) = self.into_connection_and_future(transport, async move |connection| { - main_fn(connection_context::from_raw::(connection)).await - }); + let (_, future) = + self.into_connection_and_future(transport, false, async move |connection| { + main_fn(connection_context::from_raw::(connection)).await + }); future.await } @@ -1891,6 +1892,7 @@ impl< fn into_connection_and_future( self, transport: impl ConnectTo + 'static, + wait_owned_transport: bool, main_fn: impl AsyncFnOnce(ConnectionTo) -> Result, ) -> ( ConnectionTo, @@ -1909,11 +1911,18 @@ impl< 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 (foreground_succeeded_tx, foreground_succeeded) = completion_signal(); + let (foreground_done_tx, foreground_done) = completion_signal(); let pending_replies = PendingReplies::default(); // Normalize the transport without losing ownership or finish metadata. let transport_component = crate::DynConnectTo::new(transport); - let (transport_channel, transport_future) = transport_component.into_channel_and_future(); + let (transport_channel, mut transport_future) = + transport_component.into_channel_and_future(); + let owned_transport = transport_future.is_some(); + let transport_finish = transport_future + .as_mut() + .and_then(crate::ConnectionDriver::take_finish); let (transport_completion_tx, transport_completion_rx) = oneshot::channel(); let transport_completion = transport_completion_rx .map(|result| { @@ -1935,61 +1944,107 @@ impl< pending_replies.registrar(), protocol_mode, ); - let spawn_result = if let Some(driver) = transport_future { - connection.spawn(async move { + // Transport progress must outlive successful foreground completion. + // Application tasks remain cancellable. The inherited close-callback + // phase still polls them for cleanup, but physical drain alone does not. + let transport_driver = if let Some(driver) = transport_future { + async move { let result = driver.await; drop(transport_completion_tx.send(result.clone())); result - }) + } + .boxed() } else { // Channel-only endpoints have no physical sink work to await. // Their protocol drain marker still orders accepted output. drop(transport_completion_tx.send(Ok(()))); - Ok(()) + future::ready(Ok(())).boxed() }; // Destructure the channel endpoints let Channel { - rx: transport_incoming_rx, + rx: mut transport_incoming_rx, tx: transport_outgoing_tx, } = transport_channel; + let transport_incoming = futures::stream::poll_fn({ + let mut completion = connection.transport_completion.clone(); + let mut completed = false; + move |cx| { + if owned_transport + && !completed + && let std::task::Poll::Ready(Ok(())) = + std::pin::Pin::new(&mut completion).poll(cx) + { + // Owned completion closes the producer boundary, not the + // accepted buffer. The incoming actor still dispatches every + // accepted frame and completes its close callbacks in order. + transport_incoming_rx.close(); + completed = true; + } + transport_incoming_rx.poll_next_unpin(cx) + } + }); + let protocol_compat = ProtocolCompat::new(protocol_mode); 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(), - &connection, - transport_incoming_rx, - dynamic_handler_rx, - pending_replies.clone(), - incoming_actor::IncomingHandlers::new(handler, on_close), - protocol_compat.clone(), - ); + let incoming = { + let pending_replies = pending_replies.clone(); + let protocol_compat = protocol_compat.clone(); + async { + let mut transport_incoming = std::pin::pin!(transport_incoming); + let incoming = incoming_actor::incoming_protocol_actor( + me.counterpart(), + &connection, + transport_incoming.as_mut(), + dynamic_handler_rx, + pending_replies, + incoming_actor::IncomingHandlers::new( + handler, + on_close, + foreground_succeeded.clone(), + ), + protocol_compat, + ); + // Success stops delivery, not physical I/O. An + // underway close callback finishes before sealing. + run_incoming_until_foreground_succeeds( + incoming, + foreground_succeeded, + connection.incoming_closed.clone(), + ) + .await?; + // Keep the raw producer boundary alive while the + // physical driver drains. Discard without application + // delivery, rather than fail a late read-side send. + while transport_incoming.next().await.is_some() {} + Ok(()) + } + }; let other_actors = async { futures::try_join!( + // A ready driver error is authoritative even if its + // closed channel would also make output forwarding fail. + transport_driver, // Protocol layer: OutgoingMessage -> RawJsonRpcMessage outgoing_actor::outgoing_protocol_actor( outgoing_rx, pending_replies, transport_outgoing_tx, protocol_compat, + foreground_done, ), - task_actor::task_actor(new_task_rx, &connection), - runner.run_with_connection_to(connection.clone()), )?; Ok(()) }; - // EOF can wake a pending request consumer, which may make - // the task actor fail while close callbacks are running. - // Keep the incoming actor alive until those callbacks have - // all finished, just as we do when the foreground wakes. + // Keep close callbacks alive when another core actor fails. + // The outer coordination provides the same protection when + // EOF wakes an application task or the foreground. run_until_connection_close( incoming, other_actors, @@ -2000,7 +2055,38 @@ impl< run_until_connection_close( background, - main_fn(connection.clone()), + async { + let application = async { + futures::try_join!( + task_actor::task_actor(new_task_rx, &connection), + runner.run_with_connection_to(connection.clone()), + )?; + Ok(()) + }; + let result = run_until_connection_close( + application, + async { + let result = main_fn(connection.clone()).await; + if result.is_ok() { + // Stop new incoming delivery immediately, + // including during the callback cleanup phase. + let _ = foreground_succeeded_tx.send(()); + } + // Shutdown cancels local consumers, not remote + // requests. Do not add cancellation traffic while + // dropping those consumers before the drain. + connection.pending_replies.disarm_cancellations(); + result + }, + connection.incoming_closed.clone(), + ) + .await?; + let _ = foreground_done_tx.send(()); + connection + .drain_outgoing(transport_finish, wait_owned_transport) + .await?; + Ok(result) + }, connection.incoming_closed.clone(), ) .await @@ -2262,6 +2348,15 @@ struct PendingRepliesRegistrar { } impl PendingRepliesRegistrar { + fn disarm_cancellations(&self) { + if let Some(inner) = self.inner.upgrade() { + let inner = inner.lock().expect("pending replies mutex poisoned"); + for reply in inner.replies.values() { + reply.cancellation_disarm.disarm(); + } + } + } + /// Register a response destination before the request becomes observable. /// /// Returns `false` after failing `reply` when EOF has already made a @@ -3461,6 +3556,21 @@ pub struct ConnectionTo { type SharedTransportCompletion = future::Shared>>; +type SharedCompletionSignal = future::Shared>; + +fn completion_signal() -> (oneshot::Sender<()>, SharedCompletionSignal) { + let (tx, rx) = oneshot::channel(); + let signal = async move { + // Dropping a sender (e.g. foreground failure) is not success. + if rx.await.is_err() { + future::pending::<()>().await; + } + } + .boxed() + .shared(); + (tx, signal) +} + #[derive(Clone)] struct IncomingClosed { state: Arc, @@ -3553,6 +3663,25 @@ fn incoming_transport_closed_error(method: &str) -> crate::Error { })) } +/// Unlike the cleanup coordinator below, check success before polling delivery: +/// an already-ready success must not resume a message handler into another +/// dispatch (including another entry of the same batch). +fn run_incoming_until_foreground_succeeds( + incoming: impl Future>, + foreground_succeeded: SharedCompletionSignal, + incoming_closed: IncomingClosed, +) -> impl Future> { + let mut incoming = Box::pin(incoming); + future::poll_fn(move |cx| { + if foreground_succeeded.clone().poll_unpin(cx).is_ready() && !incoming_closed.is_closing() { + return std::task::Poll::Ready(Ok(())); + } + // A close callback already underway is protected. Its result is polled + // before stopping, so callback errors retain their existing precedence. + incoming.as_mut().poll(cx) + }) +} + /// Run the connection background alongside its foreground while ensuring that /// a foreground woken by incoming EOF cannot cancel close callbacks midway. fn run_until_connection_close( @@ -3646,9 +3775,14 @@ impl ConnectionTo { self.incoming_closed.is_closed() } - /// Stop accepting outgoing messages, drain those already accepted through - /// the protocol actor, and wait for the transport sink to finish them. - async fn drain_outgoing(&self) -> Result<(), crate::Error> { + /// Stop accepting outgoing messages, drain routable output through the + /// protocol actor, and finish cooperative physical sinks. Reactive serving also + /// joins owned transport work after incoming EOF. + async fn drain_outgoing( + &self, + finish: Option, + wait_owned_transport: bool, + ) -> Result<(), crate::Error> { let (done_tx, done_rx) = oneshot::channel(); let marker_result = send_raw_message( &self.message_tx, @@ -3663,10 +3797,20 @@ impl ConnectionTo { Err(error) => Err(error), }; - // The marker only proves that all accepted protocol messages entered - // the raw transport queue. Transport completion is the sink-level - // barrier that proves a backpressured writer finished them. - self.transport_completion.clone().await?; + let physical_finish = finish.is_some(); + if let Some(mut finish) = finish { + // Only finish the physical sink after the protocol actor has handed + // off its accepted output. Closing it earlier races the drain. + finish.request(); + } + if physical_finish || wait_owned_transport { + // Cooperative completion proves physical sink drain. Reactive serving + // also waits for owned work after EOF (e.g. child exit status). + self.transport_completion.clone().await?; + } + // Opaque application drivers have no physical finish contract. Keep + // polling their errors in the background, but do not globally join work + // which may intentionally run forever after the foreground returns. marker_result } @@ -3840,7 +3984,7 @@ impl ConnectionTo { transport: impl ConnectTo + 'static, ) -> Result, crate::Error> { let (connection, future) = - builder.into_connection_and_future(transport, |_| std::future::pending()); + builder.into_connection_and_future(transport, false, |_| std::future::pending()); Task::new(std::panic::Location::caller(), future).spawn(&self.task_tx)?; Ok(connection) } @@ -6331,7 +6475,9 @@ where Either::Right((result, _)) => result, } }, - finish_tx, + move || { + let _ = finish_tx.send(()); + }, ); (channel_for_caller, server_future) @@ -6345,7 +6491,7 @@ where { async fn connect_to(self, client: impl ConnectTo) -> Result<(), crate::Error> { let (channel, mut serve_self) = self.into_channel_transport(); - let finish = serve_self + let mut finish = serve_self .take_finish() .expect("built-in Lines transport supports explicit finishing"); let client_future = Box::pin(ConnectTo::::connect_to(channel, client)); @@ -6355,7 +6501,7 @@ where result?; // The local bridge has transferred all accepted client output. // Finish the physical sink without waiting for remote read EOF. - let _ = finish.send(()); + finish.request(); serve_self.await } Either::Right((result, _)) => result, @@ -6433,16 +6579,13 @@ where let Self { outgoing, incoming } = self; let incoming_lines = Box::pin(BufReader::new(incoming).lines()); - let outgoing_lines = - futures::sink::unfold(Box::pin(outgoing), async move |mut writer, line: String| { - write_line(&mut writer, line).await?; - Ok::<_, std::io::Error>(writer) - }); + let outgoing_lines = transport_actor::LineWriter::new(outgoing); Lines::new(outgoing_lines, incoming_lines) } } +#[cfg(any(not(target_family = "wasm"), test))] pub(crate) async fn write_line(writer: &mut W, line: String) -> std::io::Result<()> where W: AsyncWrite + Unpin + ?Sized, @@ -6528,47 +6671,61 @@ impl Channel { /// Copy output concurrently with its owning driver, then drain accepted frames. /// Passive endpoints instead retain the channel's independent half-close lifetime. pub(crate) async fn copy_with_driver( - mut self, + self, driver: Option, ) -> Result<(), crate::Error> { - let Some(mut driver) = driver else { - return self.copy().await; - }; + self.copy_with_driver_until(driver, future::pending()).await + } + /// After the destination's owned foreground finishes, keep driving source + /// errors and sink work, but never deliver queued or new input to it. + pub(crate) async fn copy_with_driver_until( + mut self, + mut driver: Option, + stop_delivery: impl Future, + ) -> Result<(), crate::Error> { + let mut stop_delivery = pin!(stop_delivery); + let mut delivering = true; let mut done = false; loop { - let frame = if done { - self.rx.next().await - } else { - // Driver errors remain authoritative even when output EOF is ready. - let event = future::poll_fn(|cx| { - if let std::task::Poll::Ready(result) = std::pin::Pin::new(&mut driver).poll(cx) - { - return std::task::Poll::Ready(Either::Left(result)); - } - self.rx.poll_next_unpin(cx).map(Either::Right) - }) - .await; - match event { - Either::Left(result) => { - result?; - done = true; - self.rx.close(); - continue; - } - Either::Right(frame) => frame, + let event = future::poll_fn(|cx| { + if delivering && stop_delivery.as_mut().poll(cx).is_ready() { + delivering = false; + } + // Driver errors remain authoritative even when stop or EOF is ready. + if !done + && let Some(driver) = driver.as_mut() + && let std::task::Poll::Ready(result) = std::pin::Pin::new(driver).poll(cx) + { + return std::task::Poll::Ready(Either::Left(result)); + } + if !delivering && driver.is_none() { + return std::task::Poll::Ready(Either::Right(None)); + } + self.rx.poll_next_unpin(cx).map(Either::Right) + }) + .await; + let frame = match event { + Either::Left(result) => { + result?; + done = true; + self.rx.close(); + continue; } + Either::Right(frame) => frame, }; let Some(frame) = frame else { break; }; - self.tx - .unbounded_send(frame) - .map_err(crate::util::internal_error)?; + if delivering { + self.tx + .unbounded_send(frame) + .map_err(crate::util::internal_error)?; + } } // Propagate this half-close before waiting for a still-running driver. drop(self); - if !done { + if !done && let Some(driver) = driver { driver.await?; } Ok(()) @@ -6700,7 +6857,7 @@ mod tests { }); let incoming = futures::stream::iter([Err(std::io::Error::other("finish read failed"))]); let (_channel, mut driver) = Lines::new(outgoing, incoming).into_channel_transport(); - driver.take_finish().unwrap().send(()).unwrap(); + assert!(driver.request_finish()); let error = futures::executor::block_on(driver).unwrap_err(); assert_eq!( @@ -6873,6 +7030,7 @@ mod tests { pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::v2_proxy()), + future::pending::<()>().boxed().shared(), )); assert!( actor.as_mut().now_or_never().is_none(), @@ -7207,6 +7365,7 @@ mod tests { pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::disabled()), + future::pending::<()>().boxed().shared(), )); assert!( @@ -7238,6 +7397,93 @@ mod tests { drop(sent); } + #[test] + fn foreground_finish_settles_unready_requests_and_preserves_ready_output_fifo() { + let (connection, message_rx, pending_replies) = connection_for_response_hook_tests(); + let unready = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("unready", serde_json::json!({})).unwrap(), + future::pending(), + ); + let unready_id = unready.id().clone(); + send_raw_message( + &connection.message_tx, + OutgoingMessage::Notification { + untyped: UntypedMessage::new("first", serde_json::json!({})).unwrap(), + }, + ) + .unwrap(); + let ready = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("ready", serde_json::json!({})).unwrap(), + future::ready(Ok(())), + ); + let unready_after = connection.send_ordered_request_to_after( + crate::role::UntypedRole, + UntypedMessage::new("unready-after", serde_json::json!({})).unwrap(), + future::pending(), + ); + let unready_after_id = unready_after.id().clone(); + send_raw_message( + &connection.message_tx, + OutgoingMessage::Notification { + untyped: UntypedMessage::new("last", serde_json::json!({})).unwrap(), + }, + ) + .unwrap(); + let (done_tx, done_rx) = oneshot::channel(); + send_raw_message( + &connection.message_tx, + OutgoingMessage::CloseAfterDraining { done: done_tx }, + ) + .unwrap(); + let (transport_tx, transport_rx) = mpsc::unbounded(); + futures::executor::block_on(outgoing_actor::outgoing_protocol_actor( + message_rx, + pending_replies.clone(), + transport_tx, + ProtocolCompat::new(ProtocolMode::disabled()), + future::ready(()).boxed().shared(), + )) + .unwrap(); + futures::executor::block_on(done_rx).unwrap(); + let error = futures::executor::block_on(unready.block_task()) + .expect_err("an unresolved gate must explicitly fail its consumer"); + assert!( + error + .data + .unwrap() + .to_string() + .contains("foreground completed before outgoing request readiness") + ); + assert!(!pending_replies.contains(&unready_id)); + let error = futures::executor::block_on(unready_after.block_task()) + .expect_err("each unresolved gate must fail without repolling a consumed signal"); + assert!( + error + .data + .unwrap() + .to_string() + .contains("foreground completed before outgoing request readiness") + ); + assert!(!pending_replies.contains(&unready_after_id)); + assert!(pending_replies.contains(ready.id())); + let frames = futures::executor::block_on(transport_rx.collect::>()); + let methods = frames + .into_iter() + .map(|frame| match frame { + TransportFrame::Single(RawJsonRpcMessage::Notification(message)) => { + message.method.to_string() + } + TransportFrame::Single(RawJsonRpcMessage::Request(message)) => { + message.method.to_string() + } + _ => panic!("expected ready request/notification output"), + }) + .collect::>(); + assert_eq!(methods, ["first", "ready", "last"]); + } + #[test] fn ordered_blocking_transform_precedes_response_acknowledgment() { let (connection, _message_rx, pending_replies) = connection_for_response_hook_tests(); @@ -7319,6 +7565,7 @@ mod tests { pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::disabled()), + future::pending::<()>().boxed().shared(), )); assert!( diff --git a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs index 521037eb..4b8656d3 100644 --- a/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/incoming_actor.rs @@ -1,4 +1,5 @@ // Types re-exported from crate root +use futures::FutureExt as _; use futures::StreamExt as _; use futures::channel::mpsc; use futures::stream; @@ -41,11 +42,20 @@ use super::Handled; pub(super) struct IncomingHandlers { messages: Message, close: Close, + foreground_succeeded: super::SharedCompletionSignal, } impl IncomingHandlers { - pub(super) fn new(messages: Message, close: Close) -> Self { - Self { messages, close } + pub(super) fn new( + messages: Message, + close: Close, + foreground_succeeded: super::SharedCompletionSignal, + ) -> Self { + Self { + messages, + close, + foreground_succeeded, + } } } @@ -60,7 +70,7 @@ impl IncomingHandlers { pub(super) async fn incoming_protocol_actor( counterpart: Counterpart, connection: &ConnectionTo, - transport_rx: mpsc::UnboundedReceiver, + transport_rx: impl futures::Stream, dynamic_handler_rx: mpsc::UnboundedReceiver>, pending_replies: PendingReplies, handlers: IncomingHandlers< @@ -72,6 +82,7 @@ pub(super) async fn incoming_protocol_actor( let IncomingHandlers { messages: mut handler, close: on_close, + foreground_succeeded, } = handlers; // `merge` does not expose when one of its source streams ends. Preserve @@ -81,8 +92,9 @@ pub(super) async fn incoming_protocol_actor( transport_rx.map(IncomingProtocolMsg::Transport), stream::iter([IncomingProtocolMsg::TransportClosed]), ); - let mut my_rx = + let my_rx = transport_with_close.merge(dynamic_handler_rx.map(IncomingProtocolMsg::DynamicHandler)); + let mut my_rx = std::pin::pin!(my_rx); let mut dynamic_handlers: FxHashMap>> = FxHashMap::default(); @@ -101,6 +113,13 @@ pub(super) async fn incoming_protocol_actor( }; message }; + // Check after the receive await as well: success may have occurred + // while this actor was waiting for its next message/EOF. The caller + // cancels a blocked message handler, but protects an underway close + // callback. Neither path may start another delivery after success. + if foreground_succeeded.clone().now_or_never().is_some() { + break; + } tracing::trace!(message = ?message_result, actor = "incoming_protocol_actor"); match message_result { IncomingProtocolMsg::TransportClosed => { diff --git a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs index effe164c..970e31d4 100644 --- a/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/outgoing_actor.rs @@ -29,6 +29,7 @@ pub(super) async fn outgoing_protocol_actor( pending_replies: PendingReplies, transport_tx: mpsc::UnboundedSender, protocol_compat: ProtocolCompat, + foreground_done: super::SharedCompletionSignal, ) -> Result<(), crate::Error> { let mut drain_waiters = Vec::new(); @@ -99,19 +100,35 @@ 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 { + // A route-installation gate may depend on an application + // handler that never returns. Foreground success cancels + // unresolved gates, but ready gates still publish in FIFO + // order; shutdown must never bypass route readiness. + // Poll a fresh clone for each gate: a Shared handle that + // returned Ready cannot itself be polled a second time. + let result = + match futures::future::select(Box::pin(readiness), foreground_done.clone()) + .await + { + futures::future::Either::Left((result, _)) => result, + futures::future::Either::Right(((), _)) => { + Err(crate::Error::internal_error() + .data("foreground completed before outgoing request readiness")) + } + }; + if let Err(error) = result { + tracing::warn!( + ?id, + %method, + ?error, + "Outgoing request readiness failed" + ); + if let Some(pending_reply) = pending_replies.remove(&id) { + pending_reply.fail(error); + } + continue; } - continue; } if !pending_replies.contains(&id) { diff --git a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs index fcbb3664..806e7da2 100644 --- a/src/agent-client-protocol/src/jsonrpc/transport_actor.rs +++ b/src/agent-client-protocol/src/jsonrpc/transport_actor.rs @@ -186,7 +186,81 @@ async fn transport_outgoing_frames_actor( } } } - Ok(()) + outgoing_lines + .close() + .await + .map_err(crate::Error::into_internal_error) +} + +/// A newline writer whose sink close reaches the actual write half. An unfold +/// sink can flush each line, but has no way to forward `AsyncWrite::close`. +pub(super) struct LineWriter { + writer: std::pin::Pin>, + bytes: Vec, + written: usize, + closing: bool, +} + +impl LineWriter { + pub(super) fn new(writer: W) -> Self { + Self { + writer: Box::pin(writer), + bytes: Vec::new(), + written: 0, + closing: false, + } + } +} + +impl futures::Sink for LineWriter { + type Error = std::io::Error; + + fn poll_ready( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.poll_flush(cx) + } + + fn start_send(self: std::pin::Pin<&mut Self>, line: String) -> Result<(), Self::Error> { + let this = self.get_mut(); + this.bytes = line.into_bytes(); + this.bytes.push(b'\n'); + this.written = 0; + Ok(()) + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + let this = self.get_mut(); + while this.written < this.bytes.len() { + let count = futures::ready!( + this.writer + .as_mut() + .poll_write(cx, &this.bytes[this.written..]) + )?; + if count == 0 { + return std::task::Poll::Ready(Err(std::io::ErrorKind::WriteZero.into())); + } + this.written += count; + } + this.bytes.clear(); + this.written = 0; + this.writer.as_mut().poll_flush(cx) + } + + fn poll_close( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + if !self.closing { + futures::ready!(self.as_mut().poll_flush(cx))?; + self.closing = true; + } + self.get_mut().writer.as_mut().poll_close(cx) + } } fn malformed_line_value(raw: String) -> Result { @@ -257,6 +331,56 @@ mod tests { use super::*; use crate::ErrorCode; + #[derive(Default)] + struct PendingCloseWriter { + bytes: Vec, + close_polls: usize, + } + + impl futures::AsyncWrite for PendingCloseWriter { + fn poll_write( + mut self: std::pin::Pin<&mut Self>, + _: &mut std::task::Context<'_>, + bytes: &[u8], + ) -> std::task::Poll> { + self.bytes.extend_from_slice(bytes); + std::task::Poll::Ready(Ok(bytes.len())) + } + + fn poll_flush( + self: std::pin::Pin<&mut Self>, + _: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + assert_eq!(self.close_polls, 0, "do not flush again during shutdown"); + std::task::Poll::Ready(Ok(())) + } + + fn poll_close( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + self.close_polls += 1; + if self.close_polls == 1 { + cx.waker().wake_by_ref(); + std::task::Poll::Pending + } else { + std::task::Poll::Ready(Ok(())) + } + } + } + + #[test] + fn byte_writer_flushes_then_continues_pending_shutdown_without_reflushing() { + use futures::SinkExt as _; + let mut sink = LineWriter::new(PendingCloseWriter::default()); + futures::executor::block_on(async { + sink.send("line".to_string()).await.unwrap(); + sink.close().await.unwrap(); + }); + assert_eq!(sink.writer.bytes, b"line\n"); + assert_eq!(sink.writer.close_polls, 2); + } + #[test] fn parses_batch_entries_independently() { let ParsedIncomingLine::Batch(batch) = parse_incoming_line( diff --git a/src/agent-client-protocol/src/role/acp.rs b/src/agent-client-protocol/src/role/acp.rs index 14e48384..cb504821 100644 --- a/src/agent-client-protocol/src/role/acp.rs +++ b/src/agent-client-protocol/src/role/acp.rs @@ -439,7 +439,7 @@ impl ConnectTo for AgentProtocolRouter { let agent = RunningProtocolPeer::new(agent); agent.send_frame(first_frame)?; - pipe_protocol_peers_until_closed(client, agent).await + pipe_protocol_peers_until_done(client, agent).await } } @@ -958,10 +958,20 @@ async fn reject_initialize( error: crate::Error, ) -> Result<(), crate::Error> { let RunningProtocolPeer { mut rx, tx, driver } = client; - let future = driver.into_driver(); send_initialize_error(&tx, frame, error)?; drop(tx); + let Some(mut driver) = driver.into_driver() else { + // The rejection has already been handed to the raw channel. There is + // no owned transport work or physical drain to await. + return Ok(()); + }; + if !driver.request_finish() { + // An opaque driver has no finite physical-finish contract. Preserve a + // ready error before cancelling it rather than wait for remote EOF. + return crate::util::run_until(driver, future::ready(Ok(()))).await; + } + let drain_incoming = async move { // Later input has no protocol meaning once initialization is rejected. // Keep draining it only so the transport can flush the queued rejection; @@ -970,15 +980,11 @@ async fn reject_initialize( Ok::<_, crate::Error>(()) }; - let Some(future) = future else { - return drain_incoming.await; - }; - - match future::select(future, Box::pin(drain_incoming)).await { + match future::select(driver, Box::pin(drain_incoming)).await { future::Either::Left((result, _)) => result, - future::Either::Right((result, future)) => { + future::Either::Right((result, driver)) => { result?; - future.await + driver.await } } } @@ -995,7 +1001,7 @@ enum ProtocolPeerDriver { Passive, Active(crate::ConnectionDriver), Completed { - finish: Option>, + finish: Option, }, } @@ -1009,7 +1015,11 @@ impl ProtocolPeerDriver { // bridge, never while reading its remaining queued frames. This // records actual owned completion, not a passive ready sentinel. Self::Completed { finish } => Some(match finish { - Some(finish) => crate::ConnectionDriver::with_finish(future::ready(Ok(())), finish), + Some(mut finish) => { + crate::ConnectionDriver::with_finish(future::ready(Ok(())), move || { + finish.request(); + }) + } None => crate::ConnectionDriver::new(future::ready(Ok(()))), }), } @@ -1126,27 +1136,9 @@ fn initialize_message_mut( } } -#[cfg(feature = "unstable_protocol_v2")] -async fn pipe_protocol_peers_until_closed( - left: RunningProtocolPeer, - right: RunningProtocolPeer, -) -> Result<(), crate::Error> { - let ((), ()) = futures::try_join!( - Channel { - rx: left.rx, - tx: right.tx, - } - .copy_with_driver(left.driver.into_driver()), - Channel { - rx: right.rx, - tx: left.tx, - } - .copy_with_driver(right.driver.into_driver()), - )?; - - Ok(()) -} - +// Every protocol router uses the same ownership rule. Passive halves keep +// independent lifetimes; owned completion drains output and then either joins +// an opposed cooperative driver or cancels opaque work after polling errors. #[cfg(feature = "unstable_protocol_v2")] async fn pipe_protocol_peers_until_done( left: RunningProtocolPeer, @@ -1162,40 +1154,53 @@ async fn pipe_protocol_peers_until_done( let right_finish = right_driver .as_mut() .and_then(crate::ConnectionDriver::take_finish); + let (stop_left_tx, stop_left_rx) = futures::channel::oneshot::channel(); + let (stop_right_tx, stop_right_rx) = futures::channel::oneshot::channel(); + let stop = async |rx: futures::channel::oneshot::Receiver<()>| { + if rx.await.is_err() { + future::pending::<()>().await; + } + }; let left_to_right = Box::pin( Channel { rx: left.rx, tx: right.tx, } - .copy_with_driver(left_driver), + .copy_with_driver_until(left_driver, stop(stop_left_rx)), ); let right_to_left = Box::pin( Channel { rx: right.rx, tx: left.tx, } - .copy_with_driver(right_driver), + .copy_with_driver_until(right_driver, stop(stop_right_rx)), ); match future::select(left_to_right, right_to_left).await { future::Either::Left((result, right_to_left)) => { result?; - if left_passive || !right_passive { - if !left_passive && let Some(finish) = right_finish { - let _ = finish.send(()); + if !left_passive { + let _ = stop_right_tx.send(()); + } + if left_passive || right_finish.is_some() { + if !left_passive && let Some(mut finish) = right_finish { + finish.request(); } right_to_left.await } else { - // Passive input may remain independently open, but a ready - // forwarding error must still beat foreground success. + // Without a cooperative finish hook, opposed work may remain + // open indefinitely. Poll ready errors before cancelling it. crate::util::run_until(right_to_left, future::ready(Ok(()))).await } } future::Either::Right((result, left_to_right)) => { result?; - if right_passive || !left_passive { - if !right_passive && let Some(finish) = left_finish { - let _ = finish.send(()); + if !right_passive { + let _ = stop_left_tx.send(()); + } + if right_passive || left_finish.is_some() { + if !right_passive && let Some(mut finish) = left_finish { + finish.request(); } left_to_right.await } else { @@ -1555,7 +1560,7 @@ impl ConnectTo for ProxyProtocolRouter { let proxy = RunningProtocolPeer::new(proxy); proxy.send_frame(first_frame)?; - pipe_protocol_peers_until_closed(conductor, proxy).await + pipe_protocol_peers_until_done(conductor, proxy).await } } @@ -1846,7 +1851,9 @@ mod lifetime_tests { done_rx.await.unwrap(); Ok(()) }, - finish_tx, + move || { + let _ = finish_tx.send(()); + }, )), }; @@ -1865,7 +1872,7 @@ mod lifetime_tests { .driver .into_driver() .expect("completed owned work must not become passive"); - driver.take_finish().unwrap().send(()).unwrap(); + assert!(driver.request_finish()); finish_rx.await.unwrap(); driver.await.unwrap(); } @@ -1996,6 +2003,236 @@ mod lifetime_tests { drop(remote_input); } + #[derive(Default, Debug)] + struct GatedLineSinkState { + pending: Vec, + flushed: Vec, + closed: bool, + dropped: bool, + } + + struct GatedLineSink { + state: std::sync::Arc>, + release: futures::channel::oneshot::Receiver<()>, + released: bool, + } + + impl futures::Sink for GatedLineSink { + type Error = std::io::Error; + + fn poll_ready( + self: std::pin::Pin<&mut Self>, + _: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + std::task::Poll::Ready(Ok(())) + } + + fn start_send(self: std::pin::Pin<&mut Self>, line: String) -> Result<(), Self::Error> { + self.state.lock().unwrap().pending.push(line); + Ok(()) + } + + fn poll_flush( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + if !self.released { + if std::pin::Pin::new(&mut self.release).poll(cx).is_pending() { + return std::task::Poll::Pending; + } + self.released = true; + } + let mut state = self.state.lock().unwrap(); + let pending = std::mem::take(&mut state.pending); + state.flushed.extend(pending); + std::task::Poll::Ready(Ok(())) + } + + fn poll_close( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + match self.as_mut().poll_flush(cx) { + std::task::Poll::Ready(Ok(())) => { + self.state.lock().unwrap().closed = true; + std::task::Poll::Ready(Ok(())) + } + result => result, + } + } + } + + impl Drop for GatedLineSink { + fn drop(&mut self) { + self.state.lock().unwrap().dropped = true; + } + } + + #[tokio::test] + async fn foreground_completion_flushes_lines_with_already_normalized_remote_input() { + let state = std::sync::Arc::new(std::sync::Mutex::new(GatedLineSinkState::default())); + let (release_tx, release_rx) = futures::channel::oneshot::channel(); + let sink = GatedLineSink { + state: state.clone(), + release: release_rx, + released: false, + }; + let (remote_input, incoming) = futures::channel::mpsc::unbounded(); + remote_input + .unbounded_send(Ok( + r#"{"jsonrpc":"2.0","id":1,"result":{"protocolVersion":1}}"#.to_string(), + )) + .unwrap(); + remote_input + .unbounded_send(Ok( + r#"{"jsonrpc":"2.0","method":"test/queued","params":{}}"#.to_string(), + )) + .unwrap(); + let physical = RunningProtocolPeer::new::(crate::Lines::new(sink, incoming)); + // The real Lines driver reads both ready lines before this returns. + let (initialize, physical) = physical.next_frame().await.unwrap().unwrap(); + assert_eq!( + futures::Stream::size_hint(&physical.rx).0, + 1, + "trailing input must already be in the original normalized queue" + ); + + let (Channel { rx, tx }, mut local) = Channel::duplex(); + tx.unbounded_send(initialize).unwrap(); + let foreground = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(ConnectionDriver::new(async move { + assert!(local.rx.next().await.is_some(), "initialize response"); + for index in 0..3 { + local + .tx + .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::notification( + "test/final".into(), + serde_json::json!({ "index": index, "payload": "x".repeat(1024) }), + )?)) + .unwrap(); + } + // This owns the foreground input receiver: completion closes it. + drop(local); + Ok(()) + })), + }; + let mut bridge = Box::pin(pipe_protocol_peers_until_done(foreground, physical)); + let early_result = bridge.as_mut().now_or_never(); + assert!( + !state.lock().unwrap().pending.is_empty(), + "the real physical writer must accept output before the gated flush" + ); + // Never use time to establish the race: only the sink gate controls drain. + let _released = release_tx.send(()); + let result = match early_result { + Some(result) => result, + None => tokio::time::timeout(std::time::Duration::from_secs(1), bridge) + .await + .expect("physical drain must not require remote input EOF"), + }; + let state = state.lock().unwrap(); + assert!( + result.is_ok() && state.flushed.len() == 3 && state.closed, + "accepted output must drain cleanly despite queued remote input: result={result:?}, pending={}, flushed={}, closed={}, dropped={}", + state.pending.len(), + state.flushed.len(), + state.closed, + state.dropped, + ); + for (index, line) in state.flushed.iter().enumerate() { + let value: serde_json::Value = serde_json::from_str(line).unwrap(); + assert_eq!(value["method"], "test/final"); + assert_eq!(value["params"]["index"], index); + assert_eq!(value["params"]["payload"].as_str().unwrap().len(), 1024); + } + drop(remote_input); + } + + #[tokio::test] + async fn foreground_completion_keeps_read_errors_during_lines_drain() { + let state = std::sync::Arc::new(std::sync::Mutex::new(GatedLineSinkState::default())); + let (_release_tx, release_rx) = futures::channel::oneshot::channel(); + let sink = GatedLineSink { + state: state.clone(), + release: release_rx, + released: false, + }; + let (remote_input, incoming) = futures::channel::mpsc::unbounded(); + let physical = RunningProtocolPeer::new::(crate::Lines::new(sink, incoming)); + let (Channel { rx, tx }, local) = Channel::duplex(); + local.tx.unbounded_send(frame()).unwrap(); + drop(local); + let foreground = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))), + }; + let mut bridge = Box::pin(pipe_protocol_peers_until_done(foreground, physical)); + assert!(bridge.as_mut().now_or_never().is_none()); + assert_eq!(state.lock().unwrap().pending.len(), 1); + + // Newly read successful input is irrelevant to the completed foreground, + // but a genuine read failure must still cancel the blocked sink drain. + remote_input + .unbounded_send(Ok(r#"{"jsonrpc":"2.0","method":"late"}"#.to_string())) + .unwrap(); + remote_input + .unbounded_send(Err(std::io::Error::other("read failed after foreground"))) + .unwrap(); + let error = tokio::time::timeout(std::time::Duration::from_secs(1), bridge) + .await + .expect("read failure must not wait for the sink gate") + .expect_err("real read error must win over foreground success"); + assert!( + error + .data + .unwrap() + .to_string() + .contains("read failed after foreground"), + "the read error must not become a receiver-gone forwarding error" + ); + assert_eq!(state.lock().unwrap().flushed, Vec::::new()); + } + + #[tokio::test] + async fn foreground_completion_cancels_opposed_opaque_work_in_both_directions() { + for foreground_on_left in [true, false] { + let (Channel { rx, tx }, foreground_remote) = Channel::duplex(); + foreground_remote.tx.unbounded_send(frame()).unwrap(); + let foreground = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(ConnectionDriver::new(future::ready(Ok(())))), + }; + let (Channel { rx, tx }, mut opposed_remote) = Channel::duplex(); + let (work_tx, work_rx) = futures::channel::oneshot::channel::<()>(); + let opposed = RunningProtocolPeer { + rx, + tx, + driver: ProtocolPeerDriver::Active(ConnectionDriver::new(async move { + work_rx.await.map_err(crate::util::internal_error)?; + Ok(()) + })), + }; + let mut bridge = Box::pin(if foreground_on_left { + pipe_protocol_peers_until_done(foreground, opposed) + } else { + pipe_protocol_peers_until_done(opposed, foreground) + }); + + assert_eq!( + bridge.as_mut().now_or_never(), + Some(Ok(())), + "finite foreground must not join an opaque pending peer" + ); + assert!(work_tx.is_canceled(), "opaque work must be dropped"); + assert!(opposed_remote.rx.next().await.is_some()); + assert!(opposed_remote.rx.next().await.is_none()); + } + } + #[tokio::test] async fn foreground_completion_does_not_hide_opposed_ready_driver_error() { let (Channel { rx, tx }, _foreground_remote) = Channel::duplex(); diff --git a/src/agent-client-protocol/tests/connection_driver_normalization.rs b/src/agent-client-protocol/tests/connection_driver_normalization.rs new file mode 100644 index 00000000..8608b70e --- /dev/null +++ b/src/agent-client-protocol/tests/connection_driver_normalization.rs @@ -0,0 +1,720 @@ +//! Public-API probes for owned completion and physical transport normalization. + +use std::{ + future, io, + pin::Pin, + sync::{ + Arc, Mutex, + atomic::{AtomicUsize, Ordering}, + }, + task::{Context, Poll}, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, ByteStreams, Channel, Client, ConnectTo, ConnectionDriver, Error, JsonRpcMessage, + JsonRpcNotification, Lines, RawJsonRpcMessage, RawJsonRpcResponse, TransportBatch, + TransportFrame, UntypedMessage, + role::{Role, UntypedRole}, + schema::v1::RequestId, +}; +use futures::{FutureExt as _, StreamExt as _, future::Either}; +use serde::{Deserialize, Serialize}; +use tokio::io::{AsyncBufReadExt as _, AsyncReadExt as _, AsyncWrite, AsyncWriteExt as _}; +use tokio_util::compat::{TokioAsyncReadCompatExt as _, TokioAsyncWriteCompatExt as _}; + +const TIMEOUT: Duration = Duration::from_secs(2); + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct ProbeNotification { + sequence: usize, + payload: String, +} + +impl JsonRpcMessage for ProbeNotification { + fn matches_method(method: &str) -> bool { + method == "normalization-probe" + } + + fn method(&self) -> &'static str { + "normalization-probe" + } + + fn to_untyped_message(&self) -> Result { + UntypedMessage::new(self.method(), self) + } + + fn parse_message(method: &str, params: &impl Serialize) -> Result { + if !Self::matches_method(method) { + return Err(Error::method_not_found()); + } + agent_client_protocol::util::json_cast(params) + } +} + +impl JsonRpcNotification for ProbeNotification {} + +fn notification(sequence: usize) -> ProbeNotification { + ProbeNotification { + sequence, + payload: "x".repeat(1024), + } +} + +fn frame(sequence: usize) -> TransportFrame { + let notification = notification(sequence); + TransportFrame::Single( + RawJsonRpcMessage::notification( + notification.method().into(), + serde_json::to_value(notification).unwrap(), + ) + .unwrap(), + ) +} + +fn wire_bytes(sequence: usize) -> Vec { + let TransportFrame::Single(message) = frame(sequence) else { + unreachable!() + }; + let mut bytes = serde_json::to_vec(&message).unwrap(); + bytes.push(b'\n'); + bytes +} + +/// Counts bytes accepted by the actual duplex write half, not by an SDK queue. +/// The duplex has capacity one, so delivery requires concurrent peer reads. +struct PhysicalWriter { + inner: W, + written: Arc, +} + +impl AsyncWrite for PhysicalWriter { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + bytes: &[u8], + ) -> Poll> { + let result = Pin::new(&mut self.inner).poll_write(cx, bytes); + if let Poll::Ready(Ok(count)) = result { + self.written.fetch_add(count, Ordering::SeqCst); + } + result + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_flush(cx) + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + Pin::new(&mut self.inner).poll_shutdown(cx) + } +} + +#[tokio::test] +async fn agent_connect_with_drains_final_notification_to_physical_bytes() { + let (sdk_outgoing, mut peer_incoming) = tokio::io::duplex(1); + let (remote_input_guard, sdk_incoming) = tokio::io::duplex(1); + let written = Arc::new(AtomicUsize::new(0)); + let transport = ByteStreams::new( + PhysicalWriter { + inner: sdk_outgoing, + written: written.clone(), + } + .compat_write(), + sdk_incoming.compat(), + ); + let expected = wire_bytes(0); + let mut connection = Box::pin(async { + Agent + .builder() + .connect_with(transport, async |cx| { + cx.send_notification(notification(0))?; + Ok(()) + }) + .await + .expect("finite foreground should complete cleanly"); + assert_eq!( + written.load(Ordering::SeqCst), + expected.len(), + "connect_with returned before accepted notification reached the physical writer" + ); + }); + let mut peer = Box::pin(async { + let mut received = vec![0; expected.len()]; + peer_incoming.read_exact(&mut received).await.unwrap(); + assert_eq!(received, expected); + }); + // Borrow both futures: even a timeout keeps transport work and the peer + // read half alive until the assertion, alongside the remote input guard. + tokio::time::timeout(TIMEOUT, async { + tokio::join!(connection.as_mut(), peer.as_mut()); + }) + .await + .expect("finite foreground waited for remote EOF instead of draining accepted bytes"); + // Keep physical input open through completion and all assertions. + drop(remote_input_guard); +} + +#[tokio::test] +async fn client_connect_with_drains_final_notification_to_physical_lines() { + let (sdk_outgoing, mut peer_incoming) = tokio::io::duplex(1); + let (remote_input_guard, sdk_incoming) = tokio::io::duplex(1); + let written = Arc::new(AtomicUsize::new(0)); + let writer = PhysicalWriter { + inner: sdk_outgoing, + written: written.clone(), + }; + let outgoing = futures::sink::unfold(writer, async |mut writer, line: String| { + writer.write_all(line.as_bytes()).await?; + writer.write_all(b"\n").await?; + writer.flush().await?; + Ok::<_, io::Error>(writer) + }); + let incoming = + futures::io::AsyncBufReadExt::lines(futures::io::BufReader::new(sdk_incoming.compat())); + let transport = Lines::new(outgoing, incoming); + let expected = wire_bytes(0); + let mut connection = Box::pin(async { + Client + .builder() + .connect_with(transport, async |cx| { + cx.send_notification(notification(0))?; + Ok(()) + }) + .await + .expect("finite foreground should complete cleanly"); + assert_eq!( + written.load(Ordering::SeqCst), + expected.len(), + "connect_with returned before accepted notification reached the physical line sink" + ); + }); + let mut peer = Box::pin(async { + let mut received = vec![0; expected.len()]; + peer_incoming.read_exact(&mut received).await.unwrap(); + assert_eq!(received, expected); + }); + tokio::time::timeout(TIMEOUT, async { + tokio::join!(connection.as_mut(), peer.as_mut()); + }) + .await + .expect("finite foreground waited for remote EOF instead of draining accepted lines"); + drop(remote_input_guard); +} + +async fn finite_builder_physical_drain_probe(read_error: bool, buffered_input: bool) { + let (incoming_tx, mut incoming_rx) = futures::channel::mpsc::unbounded(); + let (eof_tx, eof_rx) = futures::channel::oneshot::channel(); + let mut eof_tx = Some(eof_tx); + let incoming = futures::stream::poll_fn(move |cx| { + let result = incoming_rx.poll_next_unpin(cx); + if matches!(result, Poll::Ready(None | Some(Err(_)))) + && let Some(eof_tx) = eof_tx.take() + { + eof_tx.send(()).unwrap(); + } + result + }); + let (entered_tx, entered_rx) = futures::channel::oneshot::channel(); + let (release_tx, release_rx) = futures::channel::oneshot::channel(); + let delivered = Arc::new(Mutex::new(Vec::new())); + let outgoing = futures::sink::unfold( + (Some(entered_tx), Some(release_rx), delivered.clone()), + async |(mut entered, mut release, delivered), line: String| { + if let Some(entered) = entered.take() { + entered.send(()).unwrap(); + release.take().unwrap().await.unwrap(); + } + delivered.lock().unwrap().push(line); + Ok::<_, io::Error>((entered, release, delivered)) + }, + ); + let closes = Arc::new(AtomicUsize::new(0)); + let callback_closes = closes.clone(); + let messages = Arc::new(AtomicUsize::new(0)); + let callback_messages = messages.clone(); + let foreground_input = incoming_tx.clone(); + let mut connection = Box::pin( + UntypedRole + .builder() + .on_receive_notification( + async move |_notification: ProbeNotification, _cx| { + callback_messages.fetch_add(1, Ordering::SeqCst); + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ) + .on_close(async move |cx| { + callback_closes.fetch_add(1, Ordering::SeqCst); + // This is valid while serving, but would fail if a late EOF + // starts the callback after the outgoing drain marker seals. + cx.send_notification(notification(1))?; + Ok(()) + }) + .connect_with(Lines::new(outgoing, incoming), async move |cx| { + cx.send_notification(notification(0))?; + if buffered_input { + // Input is accepted just as main_fn succeeds, before the + // next protocol/physical poll. Stop delivery without + // turning a live producer into "receiver is gone". + foreground_input + .unbounded_send(Ok(frame(3).to_json().unwrap())) + .unwrap(); + } + Ok(()) + }), + ); + match tokio::time::timeout( + TIMEOUT, + futures::future::select(connection.as_mut(), entered_rx), + ) + .await + .expect("accepted output never reached the gated physical sink") + { + Either::Left((result, _)) => panic!("connection completed before sink release: {result:?}"), + Either::Right((entered, _)) => entered.unwrap(), + } + assert!(connection.as_mut().now_or_never().is_none()); + assert!(delivered.lock().unwrap().is_empty()); + + // Peer EOF/error arrives only AFTER foreground success and sink entry. + // Keep the connection and sink gate alive through every assertion. + if read_error { + incoming_tx + .unbounded_send(Err(io::Error::other("late physical read failed"))) + .unwrap(); + } + drop(incoming_tx); + if read_error { + let error = tokio::time::timeout(TIMEOUT, connection.as_mut()) + .await + .expect("physical read errors must remain driven during drain") + .expect_err("physical read errors must not be hidden by foreground success"); + assert!( + error + .data + .unwrap() + .to_string() + .contains("late physical read failed") + ); + assert_eq!(closes.load(Ordering::SeqCst), 0); + assert_eq!(messages.load(Ordering::SeqCst), 0); + assert!(delivered.lock().unwrap().is_empty()); + return; + } + match tokio::time::timeout( + TIMEOUT, + futures::future::select(connection.as_mut(), eof_rx), + ) + .await + .expect("physical input stopped progressing during drain") + { + Either::Left((result, _)) => panic!("late clean EOF aborted gated drain: {result:?}"), + Either::Right((eof, _)) => eof.unwrap(), + } + let result = connection.as_mut().now_or_never(); + assert!( + result.is_none(), + "late clean EOF aborted the gated physical drain: {result:?}" + ); + assert_eq!(closes.load(Ordering::SeqCst), 0); + assert!(delivered.lock().unwrap().is_empty()); + + release_tx.send(()).unwrap(); + tokio::time::timeout(TIMEOUT, connection.as_mut()) + .await + .expect("physical drain failed to finish after sink release") + .expect("clean EOF must not discard accepted output"); + assert_eq!( + *delivered.lock().unwrap(), + vec![frame(0).to_json().unwrap()] + ); + assert_eq!(closes.load(Ordering::SeqCst), 0); + assert_eq!(messages.load(Ordering::SeqCst), 0); +} + +#[tokio::test] +async fn finite_builder_does_not_start_close_delivery_during_physical_drain() { + finite_builder_physical_drain_probe(false, false).await; +} + +#[tokio::test] +async fn finite_builder_keeps_physical_read_errors_during_drain() { + finite_builder_physical_drain_probe(true, false).await; +} + +#[tokio::test] +async fn finite_builder_stops_buffered_input_delivery_without_closing_raw_receiver_early() { + finite_builder_physical_drain_probe(false, true).await; +} + +#[tokio::test] +async fn underway_close_keeps_callback_phase_task_progress_before_sealing_output() { + let (channel, peer) = Channel::duplex(); + let Channel { mut rx, tx } = peer; + drop(tx); + let (entered_tx, entered_rx) = futures::channel::oneshot::channel(); + let (release_tx, release_rx) = futures::channel::oneshot::channel(); + let task_started = Arc::new(AtomicUsize::new(0)); + let callback_finished = Arc::new(AtomicUsize::new(0)); + let callback_done = callback_finished.clone(); + let started = task_started.clone(); + let mut connection = Box::pin( + UntypedRole + .builder() + .on_close(async move |cx| { + cx.send_notification(notification(0))?; + entered_tx.send(()).unwrap(); + release_rx.await.unwrap(); + // This send happens AFTER foreground success. The outgoing + // boundary must still be open until this callback finishes. + cx.send_notification(notification(2))?; + callback_done.store(1, Ordering::SeqCst); + Ok(()) + }) + .connect_with(channel, async move |cx| { + entered_rx.await.unwrap(); + let task_cx = cx.clone(); + // Preserve the inherited callback-phase policy: queued tasks + // can start after main_fn succeeds to unblock close cleanup. + cx.spawn(async move { + started.fetch_add(1, Ordering::SeqCst); + task_cx.send_notification(notification(1))?; + release_tx.send(()).unwrap(); + Ok(()) + })?; + Ok(37) + }), + ); + assert_eq!( + tokio::time::timeout(TIMEOUT, connection.as_mut()) + .await + .expect("close callback lost application cleanup progress") + .expect("callback output must be accepted before drain sealing"), + 37 + ); + assert_eq!(task_started.load(Ordering::SeqCst), 1); + assert_eq!(callback_finished.load(Ordering::SeqCst), 1); + let mut frames = Vec::new(); + while let Some(frame) = rx.next().now_or_never().flatten() { + frames.push(frame.to_json().unwrap()); + } + assert_eq!( + frames, + (0..3) + .map(|sequence| frame(sequence).to_json().unwrap()) + .collect::>() + ); +} + +#[tokio::test] +async fn success_does_not_resume_message_delivery_or_start_queued_tasks_for_drain() { + let (channel, peer) = Channel::duplex(); + let Channel { mut rx, tx } = peer; + tx.unbounded_send(TransportFrame::Batch( + TransportBatch::from_messages((0..2).map(|sequence| { + let TransportFrame::Single(message) = frame(sequence) else { + unreachable!() + }; + message + })) + .unwrap(), + )) + .unwrap(); + let (entered_tx, entered_rx) = futures::channel::oneshot::channel(); + let (release_tx, release_rx) = futures::channel::oneshot::channel(); + let mut entered_tx = Some(entered_tx); + let mut release_rx = Some(release_rx); + let observed = Arc::new(Mutex::new(Vec::new())); + let handler_observed = observed.clone(); + let task_started = Arc::new(AtomicUsize::new(0)); + let started = task_started.clone(); + let mut connection = Box::pin( + UntypedRole + .builder() + .on_receive_notification( + async move |notification: ProbeNotification, _cx| { + handler_observed.lock().unwrap().push(notification.sequence); + if notification.sequence == 0 { + entered_tx.take().unwrap().send(()).unwrap(); + release_rx.take().unwrap().await.unwrap(); + } + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ) + .connect_with(channel, async move |cx| { + entered_rx.await.unwrap(); + cx.spawn(async move { + started.fetch_add(1, Ordering::SeqCst); + Ok(()) + })?; + cx.send_notification(notification(2))?; + // Both this handler and main_fn become ready together. Success + // must win before the handler can start the next batch entry. + release_tx.send(()).unwrap(); + Ok(()) + }), + ); + tokio::time::timeout(TIMEOUT, connection.as_mut()) + .await + .expect("finite success awaited unrelated application work") + .expect("accepted output must still drain"); + assert_eq!(*observed.lock().unwrap(), vec![0]); + assert_eq!(task_started.load(Ordering::SeqCst), 0); + assert_eq!( + rx.next() + .now_or_never() + .flatten() + .unwrap() + .to_json() + .unwrap(), + frame(2).to_json().unwrap() + ); + // The peer input remains open through completion; no close phase justifies + // starting the queued application task. + drop(tx); +} + +struct DrivenEndpoint { + channel: Channel, + driver: Option, +} + +impl ConnectTo for DrivenEndpoint { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + let bridge = ConnectTo::::connect_to(self.channel, client); + if let Some(driver) = self.driver { + futures::try_join!(bridge, driver)?; + } else { + bridge.await?; + } + Ok(()) + } + + fn into_channel_and_future(self) -> (Channel, Option) { + (self.channel, self.driver) + } +} + +async fn reactive_completion_probe(owned: bool) { + let (channel, peer) = Channel::duplex(); + let Channel { + rx: remote_receive_guard, + tx: escaped_producer, + } = peer; + for sequence in 0..3 { + escaped_producer.unbounded_send(frame(sequence)).unwrap(); + } + // No opaque EOF-dependent future: work is already complete and all accepted + // frames are in the channel. An escaped producer is not additional owned work. + let driver = owned.then(|| ConnectionDriver::new(future::ready(Ok(())))); + let observed = Arc::new(Mutex::new(Vec::new())); + let seen = observed.clone(); + let (entered_tx, entered_rx) = futures::channel::oneshot::channel(); + let (release_tx, release_rx) = futures::channel::oneshot::channel(); + let mut entered_tx = Some(entered_tx); + let mut release_rx = Some(release_rx); + let mut connection = Box::pin( + UntypedRole + .builder() + .on_receive_notification( + async move |notification: ProbeNotification, _cx| { + seen.lock().unwrap().push(notification.sequence); + if notification.sequence == 2 { + entered_tx.take().unwrap().send(()).unwrap(); + release_rx.take().unwrap().await.unwrap(); + } + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ) + .connect_to(DrivenEndpoint { channel, driver }), + ); + // Drive both until the final accepted dispatch is inside its handler. + let dispatch = tokio::time::timeout( + TIMEOUT, + futures::future::select(connection.as_mut(), entered_rx), + ) + .await + .expect("accepted notifications never reached the reactive handler"); + match dispatch { + Either::Left((result, _)) => { + panic!("connection finished before accepted dispatch completed: {result:?}"); + } + Either::Right((entered, _)) => entered.unwrap(), + } + assert!( + connection.as_mut().now_or_never().is_none(), + "connection returned while the final accepted handler was blocked" + ); + release_tx.send(()).unwrap(); + + if owned { + // Borrow the future through the timeout. On timeout the endpoint, SDK + // actors, escaped producer, and remote receiver are all still alive. + let result = tokio::time::timeout(TIMEOUT, connection.as_mut()).await; + assert!( + result.is_ok(), + "owned completion waited for escaped producer EOF: dispatched={:?}, producer_closed={}", + observed.lock().unwrap(), + escaped_producer.is_closed() + ); + result.unwrap().expect("owned completion should be clean"); + assert_eq!(*observed.lock().unwrap(), vec![0, 1, 2]); + assert!( + escaped_producer.unbounded_send(frame(3)).is_err(), + "owned completion still accepts output from an escaped producer" + ); + } else { + // None is absence of owned work, not completion. Poll to quiescence, + // then prove the passive producer can still supply another frame. + assert!(connection.as_mut().now_or_never().is_none()); + escaped_producer.unbounded_send(frame(3)).unwrap(); + escaped_producer.close_channel(); + tokio::time::timeout(TIMEOUT, connection.as_mut()) + .await + .expect("passive connection failed to finish after real input EOF") + .expect("passive EOF should be clean"); + assert_eq!(*observed.lock().unwrap(), vec![0, 1, 2, 3]); + } + drop((escaped_producer, remote_receive_guard)); +} + +#[tokio::test] +async fn reactive_builder_owned_completion_closes_producer_after_accepted_dispatch() { + reactive_completion_probe(true).await; +} + +#[tokio::test] +async fn reactive_builder_without_driver_preserves_passive_producer() { + reactive_completion_probe(false).await; +} + +#[tokio::test] +async fn normalized_split_duplex_half_close_allows_final_reverse_response() { + // Both SDK halves share ONE underlying stream. Dropping only its write + // wrapper cannot substitute for AsyncWrite::close while its read half lives. + let (sdk_stream, peer_stream) = tokio::io::duplex(1); + let (sdk_incoming, sdk_outgoing) = tokio::io::split(sdk_stream); + let (peer_incoming, mut peer_outgoing) = tokio::io::split(peer_stream); + let transport = ByteStreams::new(sdk_outgoing.compat_write(), sdk_incoming.compat()); + let (channel, driver) = ConnectTo::::into_channel_and_future(transport); + let mut driver = Box::pin(driver.expect("ByteStreams must own transport work")); + let Channel { mut rx, tx } = channel; + tx.unbounded_send(TransportFrame::Single( + RawJsonRpcMessage::request( + "normalization-request".into(), + serde_json::json!({}), + RequestId::Number(41), + ) + .unwrap(), + )) + .unwrap(); + drop(tx); + + let stage = Arc::new(AtomicUsize::new(0)); + let peer_stage = stage.clone(); + let mut peer = Box::pin(async move { + let mut lines = tokio::io::BufReader::new(peer_incoming).lines(); + let request = lines.next_line().await.unwrap().expect("queued request"); + let RawJsonRpcMessage::Request(request) = serde_json::from_str(&request).unwrap() else { + panic!("peer expected a request"); + }; + assert_eq!(request.id, RequestId::Number(41)); + peer_stage.store(1, Ordering::SeqCst); + assert!( + lines.next_line().await.unwrap().is_none(), + "dropping Channel.tx must produce physical write EOF" + ); + peer_stage.store(2, Ordering::SeqCst); + let response = + RawJsonRpcMessage::response(request.id, Ok(serde_json::json!({ "status": "final" }))); + let mut bytes = serde_json::to_vec(&response).unwrap(); + bytes.push(b'\n'); + peer_outgoing.write_all(&bytes).await.unwrap(); + peer_outgoing.shutdown().await.unwrap(); + peer_stage.store(3, Ordering::SeqCst); + }); + let mut receive = Box::pin(async { + let Some(TransportFrame::Single(RawJsonRpcMessage::Response(RawJsonRpcResponse::Result { + id, + result, + }))) = rx.next().await + else { + panic!("SDK read half lost the final reverse response"); + }; + assert_eq!(id, RequestId::Number(41)); + assert_eq!(result["status"], "final"); + }); + // Borrow all three futures so timeout does not first destroy either read + // half (or either peer half) and artificially create physical EOF. + let result = tokio::time::timeout(TIMEOUT, async { + let (result, (), ()) = tokio::join!(driver.as_mut(), peer.as_mut(), receive.as_mut()); + result + }) + .await; + assert!( + result.is_ok(), + "split duplex stalled: peer stage={} (0=request pending, 1=waiting for physical EOF, 2=response writing, 3=response sent); SDK driver/read half and peer halves still retained", + stage.load(Ordering::SeqCst) + ); + result + .unwrap() + .expect("normalized transport should finish cleanly"); + assert_eq!(stage.load(Ordering::SeqCst), 3); +} + +struct CloseFailWriter { + written: Arc, +} + +impl futures::AsyncWrite for CloseFailWriter { + fn poll_write( + self: Pin<&mut Self>, + _: &mut Context<'_>, + bytes: &[u8], + ) -> Poll> { + // Exercise repeated partial writes as well as close-error propagation. + let count = bytes.len().min(7); + self.written.fetch_add(count, Ordering::SeqCst); + Poll::Ready(Ok(count)) + } + + fn poll_flush(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Ready(Ok(())) + } + + fn poll_close(self: Pin<&mut Self>, _: &mut Context<'_>) -> Poll> { + Poll::Ready(Err(io::Error::other("physical close failed"))) + } +} + +#[tokio::test] +async fn direct_builder_reports_physical_close_error_after_writing_accepted_output() { + let written = Arc::new(AtomicUsize::new(0)); + let transport = ByteStreams::new( + CloseFailWriter { + written: written.clone(), + }, + futures::io::Cursor::new(Vec::::new()), + ); + let error = tokio::time::timeout( + TIMEOUT, + Agent.builder().connect_with(transport, async |cx| { + cx.send_notification(notification(0))?; + Ok(()) + }), + ) + .await + .expect("close failure must complete the connection") + .expect_err("physical close errors must not be swallowed"); + assert_eq!(written.load(Ordering::SeqCst), wire_bytes(0).len()); + assert!( + error + .data + .unwrap() + .to_string() + .contains("physical close failed") + ); +} diff --git a/src/agent-client-protocol/tests/cooperative_connection_driver.rs b/src/agent-client-protocol/tests/cooperative_connection_driver.rs new file mode 100644 index 00000000..f1eb7469 --- /dev/null +++ b/src/agent-client-protocol/tests/cooperative_connection_driver.rs @@ -0,0 +1,871 @@ +//! Public-API regressions for cooperative normalized-adapter completion. + +use std::{ + future::{self, Future}, + io, + pin::Pin, + sync::{ + Arc, + atomic::{AtomicBool, AtomicUsize, Ordering}, + }, + task::{Context, Poll}, + time::Duration, +}; + +use agent_client_protocol::{ + Channel, ConnectTo, ConnectionDriver, Error, JsonRpcMessage, JsonRpcNotification, + RawJsonRpcMessage, TransportFrame, UntypedMessage, role::UntypedRole, +}; +use futures::{ + FutureExt as _, StreamExt as _, + channel::{mpsc, oneshot}, + future::Either, + task::LocalSpawnExt as _, +}; +use serde::{Deserialize, Serialize}; +use tokio::io::{AsyncBufReadExt as _, AsyncReadExt as _, AsyncWrite, AsyncWriteExt as _}; + +const TIMEOUT: Duration = Duration::from_secs(2); +const OUTPUT_COUNT: usize = 3; +const FLUSH_ERROR: &str = "custom physical flush failed"; + +#[derive(Clone, Debug, Serialize, Deserialize)] +struct ProbeNotification { + sequence: usize, +} + +impl JsonRpcMessage for ProbeNotification { + fn matches_method(method: &str) -> bool { + method == "cooperative-driver-probe" + } + + fn method(&self) -> &'static str { + "cooperative-driver-probe" + } + + fn to_untyped_message(&self) -> Result { + UntypedMessage::new(self.method(), self) + } + + fn parse_message(method: &str, params: &impl Serialize) -> Result { + if !Self::matches_method(method) { + return Err(Error::method_not_found()); + } + agent_client_protocol::util::json_cast(params) + } +} + +impl JsonRpcNotification for ProbeNotification {} + +fn frame(sequence: usize) -> TransportFrame { + TransportFrame::Single( + RawJsonRpcMessage::notification( + "cooperative-driver-probe".into(), + serde_json::json!({ "sequence": sequence }), + ) + .unwrap(), + ) +} + +fn wire_bytes(sequence: usize) -> Vec { + let mut bytes = frame(sequence).to_json().unwrap().into_bytes(); + bytes.push(b'\n'); + bytes +} + +/// A genuine one-byte physical pipe supplies write backpressure. A separate +/// gate proves that accepting every byte is not equivalent to flushing it. +struct GatedWriter { + inner: tokio::io::DuplexStream, + write_blocked: Option>, + flush_entered: Option>, + flush_release: oneshot::Receiver<()>, + written: Arc, + shutdowns: Arc, + fail_flush: bool, +} + +impl AsyncWrite for GatedWriter { + fn poll_write( + mut self: Pin<&mut Self>, + cx: &mut Context<'_>, + bytes: &[u8], + ) -> Poll> { + let result = Pin::new(&mut self.inner).poll_write(cx, bytes); + match &result { + Poll::Ready(Ok(count)) => { + self.written.fetch_add(*count, Ordering::SeqCst); + } + Poll::Pending => { + if let Some(entered) = self.write_blocked.take() { + let _ = entered.send(()); + } + } + Poll::Ready(Err(_)) => {} + } + result + } + + fn poll_flush(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + if let Some(entered) = self.flush_entered.take() { + let _ = entered.send(()); + } + match Pin::new(&mut self.flush_release).poll(cx) { + Poll::Pending => Poll::Pending, + Poll::Ready(Err(_)) => panic!("test dropped the flush gate before release"), + Poll::Ready(Ok(())) if self.fail_flush => { + Poll::Ready(Err(io::Error::other(FLUSH_ERROR))) + } + Poll::Ready(Ok(())) => Pin::new(&mut self.inner).poll_flush(cx), + } + } + + fn poll_shutdown(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { + let result = Pin::new(&mut self.inner).poll_shutdown(cx); + if matches!(result, Poll::Ready(Ok(()))) { + self.shutdowns.fetch_add(1, Ordering::SeqCst); + } + result + } +} + +/// No private SDK actors or built-in transport normalization are involved. +struct NormalizedAdapter { + channel: Channel, + driver: ConnectionDriver, +} + +impl ConnectTo for NormalizedAdapter { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + let bridge = Box::pin(ConnectTo::::connect_to(self.channel, client)); + // The adapter owns its producers and closes them when its driver ends. + match futures::future::select(bridge, self.driver).await { + Either::Left((result, mut driver)) => { + result?; + // The bridge has handed off all accepted finite-peer output. + if driver.request_finish() { + driver.await + } else { + // Preserve an error made ready by the bridge's final poll, + // but do not await arbitrary opaque work indefinitely. + (&mut driver).now_or_never().unwrap_or(Ok(())) + } + } + Either::Right((result, _bridge)) => result, + } + } + + fn into_channel_and_future(self) -> (Channel, Option) { + (self.channel, Some(self.driver)) + } +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +enum ProbeMode { + Normalized, + Decorated, + Direct, +} + +fn observe_completion(driver: ConnectionDriver, completions: Arc) -> ConnectionDriver { + driver.map_future(move |work| { + work.inspect(move |_result| { + completions.fetch_add(1, Ordering::SeqCst); + }) + }) +} + +/// An owned finite peer, not a passive channel whose input must reach EOF. +fn finite_peer(output_count: usize, done: Option>) -> NormalizedAdapter { + let (channel, physical) = Channel::duplex(); + let driver = ConnectionDriver::new(async move { + for sequence in 0..output_count { + physical + .tx + .unbounded_send(frame(sequence)) + .map_err(Error::into_internal_error)?; + } + if let Some(done) = done { + done.send(()).unwrap(); + } + // Close every owned producer on completion; escaped adapter producers + // in the physical-drain probe remain independently alive. + drop(physical); + Ok(()) + }); + NormalizedAdapter { channel, driver } +} + +fn seal_output( + output: &mut mpsc::UnboundedReceiver, + sealed: &mut Option>, +) { + // Closing the receiver rejects escaped producers but retains accepted frames. + output.close(); + if let Some(sealed) = sealed.take() { + let _ = sealed.send(()); + } +} + +struct Harness { + adapter: NormalizedAdapter, + peer_output: tokio::io::DuplexStream, + peer_input: tokio::io::DuplexStream, + escaped_output: mpsc::UnboundedSender, + foreground_done: oneshot::Receiver<()>, + foreground_signal: oneshot::Sender<()>, + write_blocked: oneshot::Receiver<()>, + finish_requested: oneshot::Receiver<()>, + output_sealed: oneshot::Receiver<()>, + input_read: oneshot::Receiver<()>, + flush_input_read: oneshot::Receiver<()>, + flush_entered: oneshot::Receiver<()>, + flush_release: oneshot::Sender<()>, + finish_calls: Arc, + completions: Arc, + written: Arc, + shutdowns: Arc, +} + +fn harness(mode: ProbeMode, fail_flush: bool) -> Harness { + let (channel, physical) = Channel::duplex(); + let escaped_output = channel.tx.clone(); + let Channel { + rx: mut output, + tx: input, + } = physical; + let (sdk_output, peer_output) = tokio::io::duplex(1); + let (peer_input, sdk_input) = tokio::io::duplex(4096); + let (foreground_signal, foreground_done) = oneshot::channel(); + let (write_blocked_tx, write_blocked) = oneshot::channel(); + let (flush_entered_tx, flush_entered) = oneshot::channel(); + let (flush_release, flush_release_rx) = oneshot::channel(); + let (finish_tx, finish_rx) = oneshot::channel(); + let (finish_requested_tx, finish_requested) = oneshot::channel(); + let (output_sealed_tx, output_sealed) = oneshot::channel(); + let (input_read_tx, input_read) = oneshot::channel(); + let (flush_input_read_tx, flush_input_read) = oneshot::channel(); + let stop_delivery = Arc::new(AtomicBool::new(false)); + let callback_stop_delivery = stop_delivery.clone(); + let finish_calls = Arc::new(AtomicUsize::new(0)); + let callback_calls = finish_calls.clone(); + let completions = Arc::new(AtomicUsize::new(0)); + let written = Arc::new(AtomicUsize::new(0)); + let shutdowns = Arc::new(AtomicUsize::new(0)); + let mut writer = GatedWriter { + inner: sdk_output, + write_blocked: Some(write_blocked_tx), + flush_entered: Some(flush_entered_tx), + flush_release: flush_release_rx, + written: written.clone(), + shutdowns: shutdowns.clone(), + fail_flush, + }; + let outgoing = async move { + let mut finish = finish_rx.fuse(); + let mut sealed = Some(output_sealed_tx); + loop { + let next = futures::select_biased! { + signal = finish => { + // Losing the hook is not a graceful finish request. + if signal.is_ok() { + seal_output(&mut output, &mut sealed); + } + continue; + }, + next = output.next().fuse() => next, + }; + let Some(frame) = next else { break }; + let mut bytes = frame.to_json()?.into_bytes(); + bytes.push(b'\n'); + let write = writer.write_all(&bytes).fuse(); + futures::pin_mut!(write); + futures::select_biased! { + signal = finish => { + if signal.is_ok() { + seal_output(&mut output, &mut sealed); + } + // Never cancel a partially completed physical write. + write.await.map_err(Error::into_internal_error)?; + }, + result = write => result.map_err(Error::into_internal_error)?, + } + } + writer.flush().await.map_err(Error::into_internal_error)?; + writer + .shutdown() + .await + .map_err(Error::into_internal_error)?; + Ok(()) + }; + let incoming = async move { + let mut lines = tokio::io::BufReader::new(sdk_input).lines(); + let mut read = Some(input_read_tx); + let mut flush_read = Some(flush_input_read_tx); + while let Some(line) = lines + .next_line() + .await + .map_err(Error::into_internal_error)? + { + // The direct bridge drops its receiver when the finite peer ends. + // Continue reading during physical drain (including genuine read + // errors), but stop delivering successful input after finish. + // Normalized modes deliberately still forward late input to test + // the generic Builder's retain-and-discard behavior. + if !stop_delivery.load(Ordering::SeqCst) { + input + .unbounded_send(TransportFrame::parse_json(&line)) + .map_err(Error::into_internal_error)?; + } + if let Some(read) = read.take() { + let _ = read.send(()); + } else if let Some(read) = flush_read.take() { + let _ = read.send(()); + } + } + Ok::<_, Error>(()) + }; + let driver = ConnectionDriver::with_finish( + async move { + futures::pin_mut!(outgoing, incoming); + match futures::future::select(outgoing, incoming).await { + Either::Left((result, _incoming)) => result, + Either::Right((result, outgoing)) => { + result?; + outgoing.await + } + } + }, + move || { + callback_calls.fetch_add(1, Ordering::SeqCst); + if mode == ProbeMode::Direct { + callback_stop_delivery.store(true, Ordering::SeqCst); + } + let _ = finish_tx.send(()); + let _ = finish_requested_tx.send(()); + }, + ); + let driver = if mode == ProbeMode::Decorated { + observe_completion(driver, completions.clone()) + } else { + driver + }; + Harness { + adapter: NormalizedAdapter { channel, driver }, + peer_output, + peer_input, + escaped_output, + foreground_done, + foreground_signal, + write_blocked, + finish_requested, + output_sealed, + input_read, + flush_input_read, + flush_entered, + flush_release, + finish_calls, + completions, + written, + shutdowns, + } +} + +async fn observe( + connection: Pin<&mut (impl Future> + ?Sized)>, + probe: oneshot::Receiver<()>, + description: &str, +) { + match tokio::time::timeout(TIMEOUT, futures::future::select(connection, probe)) + .await + .unwrap_or_else(|_| panic!("timed out waiting for {description}")) + { + Either::Left((result, _)) => { + panic!("connection completed before {description}: {result:?}"); + } + Either::Right((signal, _)) => signal.expect("probe sender dropped"), + } +} + +async fn cooperative_drain_probe(mode: ProbeMode, fail_flush: bool) { + let Harness { + adapter, + mut peer_output, + mut peer_input, + escaped_output, + foreground_done, + foreground_signal, + write_blocked, + finish_requested, + output_sealed, + input_read, + flush_input_read, + flush_entered, + flush_release, + finish_calls, + completions, + written, + shutdowns, + } = harness(mode, fail_flush); + let delivered = Arc::new(AtomicUsize::new(0)); + let handler_delivered = delivered.clone(); + let mut connection = if mode == ProbeMode::Direct { + adapter + .connect_to(finite_peer(OUTPUT_COUNT, Some(foreground_signal))) + .map(|result| result.map(|()| 37)) + .boxed() + } else { + UntypedRole + .builder() + .on_receive_notification( + async move |_notification: ProbeNotification, _cx| { + handler_delivered.fetch_add(1, Ordering::SeqCst); + Ok(()) + }, + agent_client_protocol::on_receive_notification!(), + ) + .connect_with(adapter, async move |cx| { + for sequence in 0..OUTPUT_COUNT { + cx.send_notification(ProbeNotification { sequence })?; + } + foreground_signal.send(()).unwrap(); + Ok(37) + }) + .boxed() + }; + observe( + connection.as_mut(), + foreground_done, + "finite foreground success", + ) + .await; + observe( + connection.as_mut(), + write_blocked, + "physical write backpressure", + ) + .await; + observe( + connection.as_mut(), + finish_requested, + "cooperative finish request", + ) + .await; + observe(connection.as_mut(), output_sealed, "adapter output sealing").await; + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); + assert_eq!(completions.load(Ordering::SeqCst), 0); + assert_eq!(written.load(Ordering::SeqCst), 1); + assert_eq!(shutdowns.load(Ordering::SeqCst), 0); + assert!(escaped_output.unbounded_send(frame(99)).is_err()); + assert!(connection.as_mut().now_or_never().is_none()); + + // Input arrives only after foreground success and finish. Normalized modes + // forward it without restarting application dispatch; Direct discards it + // without cancelling the still-pending physical write/flush. + peer_input.write_all(&wire_bytes(99)).await.unwrap(); + observe( + connection.as_mut(), + input_read, + "late physical input during drain", + ) + .await; + assert_eq!(delivered.load(Ordering::SeqCst), 0); + + let expected: Vec<_> = (0..OUTPUT_COUNT).flat_map(wire_bytes).collect(); + let mut received = vec![0; expected.len()]; + { + let read = peer_output.read_exact(&mut received); + futures::pin_mut!(read); + match tokio::time::timeout(TIMEOUT, futures::future::select(connection.as_mut(), read)) + .await + .expect("physical output failed to drain after peer began reading") + { + Either::Left((result, _)) => panic!("connection bypassed flush gate: {result:?}"), + Either::Right((result, _)) => { + result.unwrap(); + } + } + } + assert_eq!( + received, expected, + "all accepted notifications must reach the pipe" + ); + observe(connection.as_mut(), flush_entered, "physical flush gate").await; + assert_eq!(written.load(Ordering::SeqCst), expected.len()); + assert!(connection.as_mut().now_or_never().is_none()); + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); + assert_eq!(completions.load(Ordering::SeqCst), 0); + assert_eq!(shutdowns.load(Ordering::SeqCst), 0); + + // A second late line forces the input loop to run while flush is blocked, + // not merely while the one-byte write pipe is backpressured. + peer_input.write_all(&wire_bytes(100)).await.unwrap(); + observe( + connection.as_mut(), + flush_input_read, + "late physical input at the flush gate", + ) + .await; + assert!(connection.as_mut().now_or_never().is_none()); + assert_eq!(completions.load(Ordering::SeqCst), 0); + flush_release.send(()).unwrap(); + let result = tokio::time::timeout(TIMEOUT, connection.as_mut()) + .await + .expect("cooperative completion waited for independently open input EOF"); + if fail_flush { + let error = result.expect_err("foreground success masked the physical flush error"); + assert_eq!( + error, + Error::into_internal_error(io::Error::other(FLUSH_ERROR)) + ); + assert_eq!(shutdowns.load(Ordering::SeqCst), 0); + } else { + assert_eq!(result.unwrap(), 37); + assert_eq!(shutdowns.load(Ordering::SeqCst), 1); + let mut extra = [0]; + assert_eq!( + tokio::time::timeout(TIMEOUT, peer_output.read(&mut extra)) + .await + .expect("adapter failed to half-close physical output") + .unwrap(), + 0 + ); + } + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); + assert_eq!( + completions.load(Ordering::SeqCst), + usize::from(mode == ProbeMode::Decorated), + "completion observation must run exactly once, on success or error" + ); + assert_eq!(delivered.load(Ordering::SeqCst), 0); + // These owners survive every assertion: neither input EOF nor dropping the + // escaped output producer is allowed to be the completion trigger. + drop((peer_input, escaped_output)); +} + +#[tokio::test] +async fn finite_builder_waits_for_custom_cooperative_drain_not_input_eof() { + cooperative_drain_probe(ProbeMode::Normalized, false).await; +} + +#[tokio::test] +async fn custom_cooperative_flush_error_overrides_foreground_success() { + cooperative_drain_probe(ProbeMode::Normalized, true).await; +} + +#[tokio::test] +async fn decorated_cooperative_driver_waits_for_physical_drain() { + cooperative_drain_probe(ProbeMode::Decorated, false).await; +} + +#[tokio::test] +async fn decorated_cooperative_driver_preserves_physical_flush_error() { + cooperative_drain_probe(ProbeMode::Decorated, true).await; +} + +#[tokio::test] +async fn direct_finite_peer_waits_for_physical_drain_not_input_eof() { + cooperative_drain_probe(ProbeMode::Direct, false).await; +} + +#[tokio::test] +async fn direct_finite_peer_preserves_physical_flush_error() { + cooperative_drain_probe(ProbeMode::Direct, true).await; +} + +#[test] +fn already_requested_finish_survives_driver_handoff() { + already_requested_finish_probe(false); +} + +#[test] +fn already_requested_finish_survives_decoration_and_driver_handoff() { + already_requested_finish_probe(true); +} + +fn already_requested_finish_probe(decorate: bool) { + let (channel, physical) = Channel::duplex(); + let (finish_tx, finish_rx) = oneshot::channel(); + let (flush_started_tx, mut flush_started_rx) = oneshot::channel(); + let (flush_release_tx, flush_release_rx) = oneshot::channel(); + let finish_calls = Arc::new(AtomicUsize::new(0)); + let callback_calls = finish_calls.clone(); + let completions = Arc::new(AtomicUsize::new(0)); + let flushed = Arc::new(AtomicBool::new(false)); + let driver_flushed = flushed.clone(); + let mut driver = ConnectionDriver::with_finish( + async move { + finish_rx.await.unwrap(); + flush_started_tx.send(()).unwrap(); + flush_release_rx.await.unwrap(); + driver_flushed.store(true, Ordering::SeqCst); + drop(physical); + Ok(()) + }, + move || { + callback_calls.fetch_add(1, Ordering::SeqCst); + finish_tx.send(()).unwrap(); + }, + ); + assert!(driver.request_finish()); + let mut driver = if decorate { + observe_completion(driver, completions.clone()) + } else { + driver + }; + if decorate { + assert!(driver.request_finish()); + } + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); + assert_eq!(completions.load(Ordering::SeqCst), 0); + + let (result_tx, mut result_rx) = oneshot::channel(); + let mut pool = futures::executor::LocalPool::new(); + pool.spawner() + .spawn_local(async move { + let result = UntypedRole + .builder() + .connect_with(NormalizedAdapter { channel, driver }, async |_cx| Ok(17)) + .await; + result_tx.send(result).unwrap(); + }) + .unwrap(); + // Drive every ready actor, not just the first Pending poll before the + // outgoing drain marker has been processed. + pool.run_until_stalled(); + assert!( + result_rx.try_recv().unwrap().is_none(), + "handoff must not downgrade already-finishing work to opaque cancellation" + ); + assert_eq!(flush_started_rx.try_recv().unwrap(), Some(())); + assert!(!flushed.load(Ordering::SeqCst)); + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); + assert_eq!(completions.load(Ordering::SeqCst), 0); + + flush_release_tx.send(()).unwrap(); + pool.run_until_stalled(); + assert_eq!(result_rx.try_recv().unwrap(), Some(Ok(17))); + assert!(flushed.load(Ordering::SeqCst)); + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); + assert_eq!(completions.load(Ordering::SeqCst), usize::from(decorate)); +} + +#[test] +fn finish_request_is_synchronous_and_one_shot_not_future_completion() { + let calls = Arc::new(AtomicUsize::new(0)); + let callback_calls = calls.clone(); + let (signal, mut received) = oneshot::channel(); + let mut driver = + ConnectionDriver::with_finish(future::pending::>(), move || { + callback_calls.fetch_add(1, Ordering::SeqCst); + signal.send(()).unwrap(); + }); + assert!(driver.request_finish()); + assert_eq!(calls.load(Ordering::SeqCst), 1); + assert_eq!((&mut received).now_or_never(), Some(Ok(()))); + assert!(driver.request_finish()); + assert!(Pin::new(&mut driver).now_or_never().is_none()); + drop(driver); + assert_eq!(calls.load(Ordering::SeqCst), 1); +} + +#[test] +fn dropping_unpolled_cooperative_driver_does_not_request_finish() { + let calls = Arc::new(AtomicUsize::new(0)); + let callback_calls = calls.clone(); + let (signal, received) = oneshot::channel(); + let driver = ConnectionDriver::with_finish(future::pending::>(), move || { + callback_calls.fetch_add(1, Ordering::SeqCst); + let _ = signal.send(()); + }); + drop(driver); + assert_eq!(calls.load(Ordering::SeqCst), 0); + assert!(matches!(received.now_or_never(), Some(Err(_)))); +} + +struct DropProbe(Arc); + +impl Drop for DropProbe { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } +} + +#[tokio::test] +async fn finite_foreground_cancels_opaque_pending_driver_instead_of_waiting() { + opaque_drain_probe(false).await; +} + +#[tokio::test] +async fn decorated_opaque_driver_stays_opaque_and_is_cancelled_without_completion() { + opaque_drain_probe(true).await; +} + +async fn opaque_drain_probe(decorate: bool) { + let (channel, mut peer) = Channel::duplex(); + let dropped = Arc::new(AtomicBool::new(false)); + let completions = Arc::new(AtomicUsize::new(0)); + let guard = DropProbe(dropped.clone()); + let mut driver = ConnectionDriver::new(async move { + future::pending::<()>().await; + drop(guard); + Ok(()) + }); + assert!( + !driver.request_finish(), + "opaque work cannot be gracefully finished" + ); + let mut driver = if decorate { + observe_completion(driver, completions.clone()) + } else { + driver + }; + assert!(!driver.request_finish()); + let mut connection = Box::pin(UntypedRole.builder().connect_with( + NormalizedAdapter { channel, driver }, + async |cx| { + cx.send_notification(ProbeNotification { sequence: 0 })?; + Ok(37) + }, + )); + assert_eq!( + tokio::time::timeout(TIMEOUT, connection.as_mut()) + .await + .expect("finite foreground waited forever for opaque work") + .unwrap(), + 37 + ); + assert!(dropped.load(Ordering::SeqCst)); + assert_eq!(completions.load(Ordering::SeqCst), 0); + assert_eq!( + peer.rx + .next() + .now_or_never() + .flatten() + .unwrap() + .to_json() + .unwrap(), + frame(0).to_json().unwrap() + ); + drop(peer); +} + +fn assert_direct_completion( + adapter: NormalizedAdapter, + peer: NormalizedAdapter, + expected: Result<(), Error>, +) { + let (result_tx, mut result_rx) = oneshot::channel(); + let mut pool = futures::executor::LocalPool::new(); + pool.spawner() + .spawn_local(async move { + result_tx.send(adapter.connect_to(peer).await).unwrap(); + }) + .unwrap(); + pool.run_until_stalled(); + assert_eq!( + result_rx.try_recv().unwrap(), + Some(expected), + "direct finite-peer completion must not await opaque pending work" + ); +} + +#[test] +fn direct_finite_peer_cancels_decorated_opaque_pending_driver() { + let (channel, physical) = Channel::duplex(); + let dropped = Arc::new(AtomicBool::new(false)); + let completions = Arc::new(AtomicUsize::new(0)); + let guard = DropProbe(dropped.clone()); + let driver = ConnectionDriver::new(async move { + future::pending::<()>().await; + drop((physical, guard)); + Ok(()) + }); + let mut driver = observe_completion(driver, completions.clone()); + assert!(!driver.request_finish()); + assert_direct_completion( + NormalizedAdapter { channel, driver }, + finite_peer(0, None), + Ok(()), + ); + assert!(dropped.load(Ordering::SeqCst)); + assert_eq!(completions.load(Ordering::SeqCst), 0); +} + +#[test] +fn direct_finite_peer_does_not_mask_a_ready_cooperative_driver_error() { + let (channel, physical) = Channel::duplex(); + let finish_calls = Arc::new(AtomicUsize::new(0)); + let callback_calls = finish_calls.clone(); + let error = Error::into_internal_error(io::Error::other("ready physical driver failed")); + let driver_error = error.clone(); + let driver = ConnectionDriver::with_finish( + async move { + drop(physical); + Err(driver_error) + }, + move || { + callback_calls.fetch_add(1, Ordering::SeqCst); + }, + ); + assert_direct_completion( + NormalizedAdapter { channel, driver }, + finite_peer(0, None), + Err(error), + ); + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); +} + +#[test] +fn direct_finite_peer_does_not_mask_an_opaque_error_readied_by_bridge_completion() { + let (channel, physical) = Channel::duplex(); + let (peer_done_tx, peer_done_rx) = oneshot::channel(); + let error = Error::into_internal_error(io::Error::other("final bridge poll readied error")); + let driver_error = error.clone(); + let driver = ConnectionDriver::new(async move { + peer_done_rx.await.unwrap(); + drop(physical); + Err(driver_error) + }); + assert_direct_completion( + NormalizedAdapter { channel, driver }, + finite_peer(0, Some(peer_done_tx)), + Err(error), + ); +} + +#[test] +fn direct_owned_driver_completion_does_not_wait_for_independent_peer() { + let (channel, physical) = Channel::duplex(); + let finish_calls = Arc::new(AtomicUsize::new(0)); + let callback_calls = finish_calls.clone(); + let driver = ConnectionDriver::with_finish( + async move { + // A real adapter must flush its accepted physical output before + // this point. This empty adapter owns/closes both producers. + drop(physical); + Ok(()) + }, + move || { + callback_calls.fetch_add(1, Ordering::SeqCst); + }, + ); + let (peer_channel, peer_physical) = Channel::duplex(); + let peer_dropped = Arc::new(AtomicBool::new(false)); + let guard = DropProbe(peer_dropped.clone()); + let peer_driver = ConnectionDriver::new(async move { + future::pending::<()>().await; + drop((peer_physical, guard)); + Ok(()) + }); + assert_direct_completion( + NormalizedAdapter { channel, driver }, + NormalizedAdapter { + channel: peer_channel, + driver: peer_driver, + }, + Ok(()), + ); + assert!(peer_dropped.load(Ordering::SeqCst)); + assert_eq!(finish_calls.load(Ordering::SeqCst), 0); +} diff --git a/src/agent-client-protocol/tests/protocol_driver_finish.rs b/src/agent-client-protocol/tests/protocol_driver_finish.rs new file mode 100644 index 00000000..d653dc2b --- /dev/null +++ b/src/agent-client-protocol/tests/protocol_driver_finish.rs @@ -0,0 +1,383 @@ +//! Protocol wrappers must preserve owned completion without requiring input EOF. + +#![cfg(feature = "unstable_protocol_v2")] + +use std::{ + future, io, + sync::{ + Arc, + atomic::{AtomicBool, AtomicUsize, Ordering}, + }, + time::Duration, +}; + +use agent_client_protocol::{ + Agent, Channel, Client, Conductor, ConnectTo, ConnectionDriver, Error, JsonRpcNotification, + Lines, Proxy, RawJsonRpcMessage, RawJsonRpcResponse, Role, TransportFrame, + schema::{InitializeProxyRequest, ProtocolVersion, v1}, +}; +use futures::{ + FutureExt as _, SinkExt as _, StreamExt as _, + channel::{mpsc, oneshot}, +}; +use serde::{Deserialize, Serialize}; +use serde_json::{Value, json}; + +const DEADLINE: Duration = Duration::from_secs(2); +const FINAL_METHOD: &str = "_test/protocol-driver-final"; + +#[derive(Debug, Clone, Serialize, Deserialize, JsonRpcNotification)] +#[notification(method = "_test/protocol-driver-final")] +struct FinalNotification { + sequence: usize, +} + +struct FiniteClient; + +impl ConnectTo for FiniteClient { + async fn connect_to(self, agent: impl ConnectTo) -> Result<(), Error> { + Client + .builder() + .connect_with(agent, async |cx| { + let response = cx + .send_request(v1::InitializeRequest::new(ProtocolVersion::V1)) + .block_task() + .await?; + assert_eq!(response.protocol_version, ProtocolVersion::V1); + cx.send_notification(FinalNotification { sequence: 1 })?; + Ok(()) + }) + .await + } +} + +struct FiniteAgent; + +impl ConnectTo for FiniteAgent { + async fn connect_to(self, client: impl ConnectTo) -> Result<(), Error> { + let (initialized, mut initialization) = mpsc::unbounded(); + Agent + .builder() + .on_receive_request( + async move |request: v1::InitializeRequest, responder, _cx| { + assert_eq!(request.protocol_version, ProtocolVersion::V1); + responder.respond(v1::InitializeResponse::new(request.protocol_version))?; + initialized + .unbounded_send(()) + .map_err(Error::into_internal_error) + }, + agent_client_protocol::on_receive_request!(), + ) + .connect_with(client, async move |cx| { + initialization.next().await.expect("initialize was handled"); + cx.send_notification(FinalNotification { sequence: 1 })?; + Ok(()) + }) + .await + } +} + +struct FiniteProxy; + +impl ConnectTo for FiniteProxy { + async fn connect_to(self, conductor: impl ConnectTo) -> Result<(), Error> { + let (initialized, mut initialization) = mpsc::unbounded(); + Proxy + .builder() + .on_receive_request_from( + Client, + async move |request: InitializeProxyRequest, responder, _cx| { + assert_eq!(request.initialize.protocol_version, ProtocolVersion::V1); + responder.respond(v1::InitializeResponse::new( + request.initialize.protocol_version, + ))?; + initialized + .unbounded_send(()) + .map_err(Error::into_internal_error) + }, + agent_client_protocol::on_receive_request!(), + ) + .connect_with(conductor, async move |cx| { + initialization.next().await.expect("initialize was handled"); + cx.send_notification_to(Client, FinalNotification { sequence: 1 })?; + Ok(()) + }) + .await + } +} + +/// Normalization must pass through the original owned driver, including its hook. +struct NormalizedAdapter { + channel: Channel, + driver: ConnectionDriver, +} + +impl ConnectTo for NormalizedAdapter { + async fn connect_to(self, peer: impl ConnectTo) -> Result<(), Error> { + futures::try_join!(ConnectTo::::connect_to(self.channel, peer), self.driver)?; + Ok(()) + } + + fn into_channel_and_future(self) -> (Channel, Option) { + (self.channel, Some(self.driver)) + } +} + +struct DropSignal(Arc); + +impl Drop for DropSignal { + fn drop(&mut self) { + self.0.store(true, Ordering::SeqCst); + } +} + +fn frame_value(frame: TransportFrame) -> Result { + serde_json::from_str(&frame.to_json()?).map_err(Error::into_internal_error) +} + +fn assert_final(value: &Value) { + assert_eq!(value["jsonrpc"], "2.0"); + assert_eq!(value["method"], FINAL_METHOD); + assert_eq!(value["params"], json!({ "sequence": 1 })); + assert!(value.get("id").is_none(), "{value}"); +} + +fn assert_rejection(value: &Value) { + assert_eq!(value["jsonrpc"], "2.0"); + assert_eq!(value["id"], 7); + assert_eq!(value["error"]["code"], -32602); + assert!(value.get("result").is_none(), "{value}"); +} + +/// Exercise Lines' actual JSON serialization and sink closure. The incoming +/// sender remains alive until both output EOF and successful router completion. +async fn run_over_lines( + component: impl ConnectTo, + method: &str, + params: Value, +) -> Result, Error> { + let (input_guard, incoming) = mpsc::unbounded::>(); + let (outgoing, mut output) = mpsc::unbounded::(); + let lines = Lines::new(outgoing.sink_map_err(io::Error::other), incoming); + input_guard + .unbounded_send(Ok(json!({ + "jsonrpc": "2.0", + "id": 7, + "method": method, + "params": params, + }) + .to_string())) + .map_err(Error::into_internal_error)?; + + let mut frames = Vec::new(); + let result = tokio::time::timeout(DEADLINE, async { + futures::try_join!(component.connect_to(lines), async { + while let Some(line) = output.next().await { + frames.push( + serde_json::from_str::(&line).map_err(Error::into_internal_error)?, + ); + } + Ok::<_, Error>(()) + }) + }) + .await; + assert!( + result.is_ok(), + "router must finish with remote input open; observed wire frames: {frames:?}" + ); + result.expect("deadline checked")?; + // Do not let input EOF make the completion assertion pass. + drop(input_guard); + Ok(frames) +} + +#[tokio::test(flavor = "current_thread")] +async fn client_connector_cancels_opaque_driver_after_final_frame() -> Result<(), Error> { + let (channel, Channel { mut rx, tx }) = Channel::duplex(); + let input_guard = tx.clone(); + let dropped = Arc::new(AtomicBool::new(false)); + let drop_signal = DropSignal(dropped.clone()); + let adapter = NormalizedAdapter { + channel, + driver: ConnectionDriver::new(async move { + // Capture the guard before polling, so cancellation is observable + // even if the opaque future has never started. + let _drop_signal = drop_signal; + future::pending::>().await + }), + }; + let final_observed = Arc::new(AtomicBool::new(false)); + let peer_observed = final_observed.clone(); + let peer = async move { + let frame = rx.next().await.expect("connector sends initialize"); + let TransportFrame::Single(RawJsonRpcMessage::Request(request)) = frame else { + panic!("expected a single initialize request, got {frame:?}"); + }; + assert_eq!(request.method.as_ref(), "initialize"); + tx.unbounded_send(TransportFrame::Single(RawJsonRpcMessage::Response( + RawJsonRpcResponse::Result { + id: request.id, + result: serde_json::to_value(v1::InitializeResponse::new(ProtocolVersion::V1)) + .map_err(Error::into_internal_error)?, + }, + ))) + .map_err(Error::into_internal_error)?; + assert_final(&frame_value( + rx.next().await.expect("final notification is handed off"), + )?); + peer_observed.store(true, Ordering::SeqCst); + Ok::<_, Error>(()) + }; + let connector = Client.protocol_connector().with_v1(|| FiniteClient); + let mut adapter = Some(adapter); + let result = tokio::time::timeout(DEADLINE, async { + futures::try_join!( + connector.connect_to(move || adapter.take().expect("v1 transport opened once")), + peer, + ) + }) + .await; + assert!( + result.is_ok(), + "connector must cancel opaque owned work after handoff; final frame observed: {}", + final_observed.load(Ordering::SeqCst) + ); + result.expect("deadline checked")?; + assert!(final_observed.load(Ordering::SeqCst)); + assert!( + dropped.load(Ordering::SeqCst), + "successful connector completion must drop the pending opaque driver" + ); + drop(input_guard); + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn rejected_initialize_flushes_lines_without_input_eof() -> Result<(), Error> { + let frames = run_over_lines( + Agent.protocol_router().with_v1(Agent.builder()), + "initialize", + json!({}), + ) + .await?; + assert_eq!(frames.len(), 1, "{frames:?}"); + assert_rejection(&frames[0]); + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn rejected_initialize_requests_custom_finish_once_after_handoff() -> Result<(), Error> { + for fail_flush in [false, true] { + let (channel, Channel { mut rx, tx }) = Channel::duplex(); + let input_guard = tx; + input_guard + .unbounded_send(TransportFrame::Single(RawJsonRpcMessage::request( + "initialize".into(), + json!({}), + v1::RequestId::Number(7), + )?)) + .map_err(Error::into_internal_error)?; + + let (finish_tx, finish_rx) = oneshot::channel(); + let (flush_started_tx, flush_started_rx) = oneshot::channel(); + let (flush_release_tx, flush_release_rx) = oneshot::channel(); + let finish_calls = Arc::new(AtomicUsize::new(0)); + let callback_calls = finish_calls.clone(); + let drained = Arc::new(AtomicBool::new(false)); + let driver_drained = drained.clone(); + let flush_error = Error::internal_error().data("rejection flush failed"); + let driver_error = flush_error.clone(); + let adapter = NormalizedAdapter { + channel, + driver: ConnectionDriver::with_finish( + async move { + // Consume no output until finish. Sealing then draining + // proves the rejection was handed off before the hook. + finish_rx + .await + .expect("explicit cooperative finish request"); + rx.close(); + let mut frames = Vec::new(); + while let Some(frame) = rx.next().await { + frames.push(frame_value(frame)?); + } + assert_eq!(frames.len(), 1, "{frames:?}"); + assert_rejection(&frames[0]); + flush_started_tx.send(()).unwrap(); + flush_release_rx.await.unwrap(); + driver_drained.store(true, Ordering::SeqCst); + if fail_flush { + Err(driver_error) + } else { + Ok(()) + } + }, + move || { + callback_calls.fetch_add(1, Ordering::SeqCst); + finish_tx.send(()).expect("owned driver is still alive"); + }, + ), + }; + let mut router = Box::pin( + Agent + .protocol_router() + .with_v1(Agent.builder()) + .connect_to(adapter), + ); + + assert!(router.as_mut().now_or_never().is_none()); + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); + assert!(matches!(flush_started_rx.now_or_never(), Some(Ok(())))); + assert!(!drained.load(Ordering::SeqCst), "flush is still gated"); + + flush_release_tx.send(()).unwrap(); + let result = tokio::time::timeout(DEADLINE, router) + .await + .expect("rejection drain must not require remote input EOF"); + if fail_flush { + assert_eq!(result, Err(flush_error)); + } else { + result?; + } + assert_eq!(finish_calls.load(Ordering::SeqCst), 1); + assert!(drained.load(Ordering::SeqCst), "owned drain was awaited"); + drop(input_guard); + } + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn finite_agent_router_flushes_response_and_final_notification() -> Result<(), Error> { + let frames = run_over_lines( + Agent.protocol_router().with_v1(FiniteAgent), + "initialize", + serde_json::to_value(v1::InitializeRequest::new(ProtocolVersion::V1)) + .map_err(Error::into_internal_error)?, + ) + .await?; + assert_eq!(frames.len(), 2, "{frames:?}"); + assert_eq!(frames[0]["id"], 7); + assert_eq!(frames[0]["result"]["protocolVersion"], 1); + assert!(frames[0].get("error").is_none(), "{frames:?}"); + assert_final(&frames[1]); + Ok(()) +} + +#[tokio::test(flavor = "current_thread")] +async fn finite_proxy_router_flushes_response_and_final_notification() -> Result<(), Error> { + let frames = run_over_lines( + Proxy.protocol_router().with_v1(FiniteProxy), + "_proxy/initialize", + serde_json::to_value(InitializeProxyRequest::from(v1::InitializeRequest::new( + ProtocolVersion::V1, + ))) + .map_err(Error::into_internal_error)?, + ) + .await?; + assert_eq!(frames.len(), 2, "{frames:?}"); + assert_eq!(frames[0]["id"], 7); + assert_eq!(frames[0]["result"]["protocolVersion"], 1); + assert!(frames[0].get("error").is_none(), "{frames:?}"); + assert_final(&frames[1]); + Ok(()) +}