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..86aaf39c --- /dev/null +++ b/md/migration-connection-drivers.md @@ -0,0 +1,221 @@ +# Migrating Connection Drivers + +`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 +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 `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 +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, Option) { + let (channel, future) = self.into_channel_transport(); + (channel, Some(ConnectionDriver::new(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, Option) { + (self.channel, None) +} +``` + +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; 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 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. +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 8204b77e..994d58dc 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,64 @@ 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, Option); ``` -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 optional driver +distinguishes owned work from a passive endpoint: + +| Returned work | Lifetime rule | +| --- | --- | +| `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 `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. +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. +`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 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 +channel types or introduce frame-size, queue, or task limits. ## Transport Implementations @@ -346,7 +394,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 40dc362a..8d2efb1e 100644 --- a/src/agent-client-protocol-http/CHANGELOG.md +++ b/src/agent-client-protocol-http/CHANGELOG.md @@ -2,11 +2,25 @@ ## [Unreleased] +### Changed + +- Adapt `HttpClient`'s `ConnectTo` conversion to the core SDK's breaking + 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. +- 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. +- 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/client.rs b/src/agent-client-protocol-http/src/client.rs index 86085762..2b4acf70 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; @@ -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, BoxFuture<'static, Result<(), AcpError>>) { + fn into_channel_and_future(self) -> (Channel, Option) { let (caller, transport) = Channel::duplex(); - (caller, Box::pin(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 f2796bab..04df3888 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, Option); } 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, Option) { 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; @@ -533,23 +540,34 @@ 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 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(); - 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; 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 +589,26 @@ impl ConnectionRegistry { } } -async fn close_connection_task(connection: Weak) { - let Some(connection) = connection.upgrade() else { - return; - }; +// 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}"); + } } - connection.close_streams(); + Some(connection) } fn pending_route_key(id: &RequestId) -> Option { @@ -608,8 +635,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,35 +767,25 @@ mod tests { } impl AgentFactory for ExitingAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, Option) { 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(()) }); - (transport, future) + (transport, Some(future)) } } struct RespondThenExitAgentFactory; impl AgentFactory for RespondThenExitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, Option) { 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( @@ -780,7 +796,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -789,15 +805,10 @@ mod tests { } impl AgentFactory for MalformedThenWaitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, Option) { 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 @@ -810,7 +821,7 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } @@ -820,16 +831,11 @@ mod tests { } impl AgentFactory for SendThenWaitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, Option) { 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)) @@ -838,7 +844,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -847,15 +853,10 @@ mod tests { } impl AgentFactory for BatchThenWaitAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, Option) { 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(), @@ -877,24 +878,22 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } 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, Option) { 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 @@ -909,10 +908,70 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } + #[tokio::test] + 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 || { + 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("absent owned work must not abort inbound forwarding") + .is_some() + ); + remote.tx.unbounded_send(frame.clone()).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()); + 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] async fn agent_exit_removes_connection_and_closes_streams() { let exit = Arc::new(Notify::new()); @@ -998,12 +1057,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 +1088,164 @@ 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" + )); + } + + 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] diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index 4f3ef608..8271186b 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, Option) { 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: _, @@ -524,21 +519,16 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } struct RejectingInitializeAgentFactory; impl AgentFactory for RejectingInitializeAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, Option) { 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 @@ -578,21 +568,16 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } struct PendingInitializeAgentFactory; impl AgentFactory for PendingInitializeAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, Option) { let (agent, transport) = Channel::duplex(); - let future = Box::pin(async move { + let future = ConnectionDriver::new(async move { let Channel { rx: mut incoming, tx: _outgoing, @@ -601,7 +586,7 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } @@ -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, Option) { 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"); }; @@ -671,21 +651,16 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } struct SideTrafficBeforeInitializeResponseAgentFactory; impl AgentFactory for SideTrafficBeforeInitializeResponseAgentFactory { - fn spawn_agent( - &self, - ) -> ( - Channel, - BoxFuture<'static, agent_client_protocol::Result<()>>, - ) { + fn spawn_agent(&self) -> (Channel, Option) { 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"); }; @@ -720,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 faf57c60..56ee7d67 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, Option) { 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, @@ -292,7 +286,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -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, Option) { 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"); }; @@ -336,7 +325,7 @@ mod tests { std::future::pending::>().await }); - (transport, future) + (transport, Some(future)) } } @@ -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, Option) { 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 @@ -368,7 +352,7 @@ mod tests { Ok(()) }); - (transport, future) + (transport, Some(future)) } } @@ -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, Option) { 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 @@ -401,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 7a3c1c68..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 @@ -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}, @@ -69,22 +69,21 @@ 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, - BoxFuture<'static, Result<(), agent_client_protocol::Error>>, - ) + fn into_channel_and_future(self) -> (Channel, Option) where Self: Sized, { let (channel_a, channel_b) = Channel::duplex(); - (channel_a, Box::pin(run(self.listener, channel_b))) + ( + channel_a, + 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 8412af36..9a7cb050 100644 --- a/src/agent-client-protocol/CHANGELOG.md +++ b/src/agent-client-protocol/CHANGELOG.md @@ -9,9 +9,24 @@ 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, 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 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 @@ -19,6 +34,35 @@ 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. +- 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 ### Added diff --git a/src/agent-client-protocol/src/acp_agent.rs b/src/agent-client-protocol/src/acp_agent.rs index db5dc7ae..e209e70a 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 { @@ -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 ba8b9297..c5bcb6bf 100644 --- a/src/agent-client-protocol/src/component.rs +++ b/src/agent-client-protocol/src/component.rs @@ -27,10 +27,202 @@ //! ``` 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}; +// 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. 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, +} + +impl ConnectionDriver { + /// Create a driver that owns the endpoint's lifetime. + /// + /// 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), + finish: None, + } + } + + /// 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: impl FnOnce() + Send + 'static, + ) -> Self { + Self { + future: Box::pin(future), + finish: Some(FinishControl::new(finish)), + } + } + + /// 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() + } +} + +impl Future for ConnectionDriver { + type Output = Result<()>; + + fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + 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("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`)). /// @@ -61,7 +253,7 @@ use crate::{Channel, Result, role::Role}; /// 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 @@ -129,7 +321,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 optional owned driver. /// /// The returned [`Channel`] is the canonical frame-aware boundary. It carries /// complete [`TransportFrame`](crate::TransportFrame) values so default @@ -137,7 +329,8 @@ pub trait ConnectTo: Send + 'static { /// /// This method returns: /// - A `Channel` that can be used to communicate with this component - /// - A `BoxFuture` that drives the component's connection logic + /// - `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. @@ -146,15 +339,44 @@ 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, 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 = Box::pin(self.connect_to(channel_b)); - (channel_a, future) + let future = ConnectionDriver::new(self.connect_to(channel_b)); + (channel_a, Some(future)) } } @@ -171,8 +393,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, Option); } /// Blanket implementation: any `ConnectTo` can be type-erased. @@ -195,9 +416,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, Option) { (*self).into_channel_and_future() } } @@ -251,7 +470,7 @@ impl ConnectTo for DynConnectTo { .await } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<()>>) { + fn into_channel_and_future(self) -> (Channel, Option) { self.inner.into_channel_and_future_erased() } } @@ -268,6 +487,175 @@ 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 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 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_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(); + 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()); + assert!( + driver.request_finish(), + "erasure must retain finish coordination" + ); + 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_none()); + } #[test] fn dyn_connect_to_reports_static_type_name_and_correct_debug_label() { diff --git a/src/agent-client-protocol/src/jsonrpc.rs b/src/agent-client-protocol/src/jsonrpc.rs index 26bed388..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,12 +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(); - // 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_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| { @@ -1936,54 +1944,107 @@ 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 - }); + // 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(()))); + 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, @@ -1994,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 @@ -2256,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 @@ -3455,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, @@ -3547,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( @@ -3640,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, @@ -3657,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 } @@ -3834,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) } @@ -6266,52 +6416,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,31 +6461,56 @@ 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, + } + }, + move || { + let _ = finish_tx.send(()); + }, + ); + + (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 mut 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. + finish.request(); + serve_self.await } Either::Right((result, _)) => result, } } - fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), crate::Error>>) { - self.into_channel_transport() + fn into_channel_and_future(self) -> (Channel, Option) { + let (channel, driver) = self.into_channel_transport(); + (channel, Some(driver)) } } @@ -6417,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, @@ -6448,7 +6607,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, Option) { ConnectTo::::into_channel_and_future(self.into_lines()) } } @@ -6509,6 +6668,69 @@ 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( + self, + driver: Option, + ) -> Result<(), crate::Error> { + 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 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; + }; + 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 && let Some(driver) = driver { + 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 +6783,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_none(); + 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, Option) { + (self, None) } } @@ -6586,6 +6821,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(); + assert!(driver.request_finish()); + + 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, @@ -6748,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(), @@ -7082,6 +7365,7 @@ mod tests { pending_replies, transport_tx, ProtocolCompat::new(ProtocolMode::disabled()), + future::pending::<()>().boxed().shared(), )); assert!( @@ -7113,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(); @@ -7194,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 62ba8f08..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 { @@ -202,7 +276,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 @@ -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/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..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 } } @@ -957,10 +957,21 @@ 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; 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; @@ -969,37 +980,94 @@ async fn reject_initialize( Ok::<_, crate::Error>(()) }; - let ((), ()) = futures::try_join!(future, drain_incoming)?; - Ok(()) + match future::select(driver, Box::pin(drain_incoming)).await { + future::Either::Left((result, _)) => result, + future::Either::Right((result, driver)) => { + result?; + driver.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 { + finish: Option, + }, +} + +#[cfg(feature = "unstable_protocol_v2")] +impl ProtocolPeerDriver { + fn into_driver(self) -> Option { + match self { + 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. This + // records actual owned completion, not a passive ready sentinel. + Self::Completed { finish } => Some(match finish { + Some(mut finish) => { + crate::ConnectionDriver::with_finish(future::ready(Ok(())), move || { + finish.request(); + }) + } + None => 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 = 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, 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(mut 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(&mut future, Box::pin(rx.next())).await { + future::Either::Right((Some(frame), _)) => Ok(Some(( + frame, + Self { + rx, + tx, + driver: ProtocolPeerDriver::Active(future), + }, + ))), + future::Either::Right((None, _)) => { + 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 +1076,9 @@ impl RunningProtocolPeer { Self { rx, tx, - future: Box::pin(future::ready(Ok(()))), + driver: ProtocolPeerDriver::Completed { + finish: future.take_finish(), + }, }, ))) } @@ -1066,56 +1136,77 @@ fn initialize_message_mut( } } +// 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_closed( +async fn pipe_protocol_peers_until_done( left: RunningProtocolPeer, right: RunningProtocolPeer, ) -> Result<(), crate::Error> { - let ((), (), (), ()) = futures::try_join!( - left.future, - right.future, + let mut left_driver = left.driver.into_driver(); + let mut right_driver = right.driver.into_driver(); + 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 (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(), + .copy_with_driver_until(left_driver, stop(stop_left_rx)), + ); + let right_to_left = Box::pin( Channel { rx: right.rx, tx: left.tx, } - .copy(), - )?; - - Ok(()) -} + .copy_with_driver_until(right_driver, stop(stop_right_rx)), + ); -#[cfg(feature = "unstable_protocol_v2")] -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, + match future::select(left_to_right, right_to_left).await { + future::Either::Left((result, right_to_left)) => { + result?; + if !left_passive { + let _ = stop_right_tx.send(()); } - .copy(), - Channel { - rx: right.rx, - tx: left.tx, + if left_passive || right_finish.is_some() { + if !left_passive && let Some(mut finish) = right_finish { + finish.request(); + } + right_to_left.await + } else { + // 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 } - .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, + } + future::Either::Right((result, left_to_right)) => { + result?; + 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 { + crate::util::run_until(left_to_right, future::ready(Ok(()))).await + } + } } } @@ -1469,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 } } @@ -1724,3 +1815,464 @@ 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 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(()) + }, + move || { + let _ = finish_tx.send(()); + }, + )), + }; + + 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"); + assert!(driver.request_finish()); + 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(); + 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(()) + }; + match driver { + Some(driver) => crate::util::run_until(driver, foreground).await, + None => 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); + } + + #[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(); + 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/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/jsonrpc_transport_close.rs b/src/agent-client-protocol/tests/jsonrpc_transport_close.rs index 8849b8a3..aca89329 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, }; @@ -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(()) } } @@ -169,6 +172,80 @@ 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); + if let Some(driver) = driver { + 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); + if let Some(driver) = driver { + driver.await?; + } + Ok(()) + } +} + +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, Option) { + (self.channel, Some(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 +354,233 @@ 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_have_no_driver() { + 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_none(), "Channel does not own runnable work"); + } +} + +#[test] +fn ready_owned_driver_is_awaitable() { + let driver = ConnectionDriver::new(future::ready(Ok(()))); + 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(); @@ -378,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_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(()) +} 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(