From ba33bdd7fbf092f32f6d29b0b7c85a321711f256 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Thu, 23 Jul 2026 15:08:00 +0200 Subject: [PATCH 1/3] fix(http): preserve session lifecycle across batches --- src/agent-client-protocol-http/src/client.rs | 1106 ++++++++++++++++-- 1 file changed, 1039 insertions(+), 67 deletions(-) diff --git a/src/agent-client-protocol-http/src/client.rs b/src/agent-client-protocol-http/src/client.rs index 5920716..0661b61 100644 --- a/src/agent-client-protocol-http/src/client.rs +++ b/src/agent-client-protocol-http/src/client.rs @@ -147,18 +147,20 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { incoming, }; let mut lifecycle = HttpTransportLifecycle::new(connection); - let mut ordered_posts = PostQueue::default(); - let mut response_posts = PostQueue::default(); + let mut posts = PostQueues::default(); + let mut buffered_outgoing = VecDeque::new(); let mut outgoing_closed = false; - let result = loop { - if outgoing_closed && ordered_posts.is_empty() && response_posts.is_empty() { + let result = 'transport: loop { + if outgoing_closed && buffered_outgoing.is_empty() && posts.is_empty() { break Ok(()); } let event = { let outgoing_next = async { - if outgoing_closed { + if let Some(frame) = buffered_outgoing.pop_front() { + Some(frame) + } else if outgoing_closed { futures::future::pending().await } else { outgoing.next().await @@ -167,8 +169,8 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { .fuse(); let sse_event_next = sse_event_rx.next().fuse(); let sse_failure_next = lifecycle.next_sse_failure().fuse(); - let ordered_post_next = ordered_posts.next_completion().fuse(); - let response_post_next = response_posts.next_completion().fuse(); + let ordered_post_next = posts.ordered.next_completion().fuse(); + let response_post_next = posts.responses.next_completion().fuse(); pin_mut!( outgoing_next, sse_event_next, @@ -198,31 +200,43 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { let Some(event) = event else { continue; }; - let open_session_id = match &event.frame { - TransportFrame::Single(message) => state.session_to_open_for_response(message), - TransportFrame::Malformed { .. } | TransportFrame::Batch(_) => None, - }; + let open_session_ids = state.sessions_to_open_for_responses(&event.frame); state.deliver_frame(event.frame); - if let Some(session_id) = open_session_id { - lifecycle.start_sse(Some(session_id), sse_event_tx.clone()); + for session_id in open_session_ids { + match lifecycle + .start_sse( + Some(session_id), + sse_event_tx.clone(), + SseStartContext { + events: &mut sse_event_rx, + outgoing: &mut outgoing, + buffered_outgoing: &mut buffered_outgoing, + posts: &mut posts, + state: &mut state, + }, + ) + .await + { + Ok(SseStartOutcome::Established) => {} + Ok(SseStartOutcome::OutgoingClosed) + if buffered_outgoing.is_empty() && posts.is_empty() => + { + break 'transport Ok(()); + } + Ok(SseStartOutcome::OutgoingClosed) => { + break 'transport Err(sse_setup_blocked_output_error()); + } + Err(error) => break 'transport Err(error), + } } continue; } HttpLoopEvent::SseFailure(failure) => { - let scope = failure.session_id.as_deref().unwrap_or("connection"); - error!(session_id = ?failure.session_id, error = %failure.error, "SSE stream ended"); - break Err(AcpError::internal_error() - .data(format!("{scope} SSE stream ended: {}", failure.error))); + break Err(sse_failure_error(failure)); } HttpLoopEvent::Post(completed) => { - let CompletedPost { - pending_request, - result, - } = completed; - if let Err(e) = result { - state.remove_pending_request(pending_request.as_ref()); - error!("POST failed: {e}"); - break Err(AcpError::internal_error().data(format!("POST: {e}"))); + if let Err(error) = handle_completed_post(&mut state, completed) { + break Err(error); } continue; } @@ -239,8 +253,35 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { match state.prepare_frame_post(frame) { // Response-only batches answer SSE-delivered callbacks and // must not be blocked behind the request they answer. - Ok(post) if is_response_only => response_posts.push(post), - Ok(post) => ordered_posts.push(post), + Ok((post, session_ids)) => { + for session_id in session_ids { + match lifecycle + .start_sse( + Some(session_id), + sse_event_tx.clone(), + SseStartContext { + events: &mut sse_event_rx, + outgoing: &mut outgoing, + buffered_outgoing: &mut buffered_outgoing, + posts: &mut posts, + state: &mut state, + }, + ) + .await + { + Ok(SseStartOutcome::Established) => {} + Ok(SseStartOutcome::OutgoingClosed) => { + break 'transport Err(sse_setup_blocked_output_error()); + } + Err(error) => break 'transport Err(error), + } + } + if is_response_only { + posts.responses.push(post); + } else { + posts.ordered.push(post); + } + } Err(error) => { error!("POST failed: {error}"); break Err(AcpError::internal_error().data(format!("POST: {error}"))); @@ -257,7 +298,29 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { } match state.initialize(msg).await { Ok(InitializeOutcome::Connected) => { - lifecycle.start_sse(None, sse_event_tx.clone()); + match lifecycle + .start_sse( + None, + sse_event_tx.clone(), + SseStartContext { + events: &mut sse_event_rx, + outgoing: &mut outgoing, + buffered_outgoing: &mut buffered_outgoing, + posts: &mut posts, + state: &mut state, + }, + ) + .await + { + Ok(SseStartOutcome::Established) => {} + Ok(SseStartOutcome::OutgoingClosed) if buffered_outgoing.is_empty() => { + break 'transport Ok(()); + } + Ok(SseStartOutcome::OutgoingClosed) => { + break 'transport Err(sse_setup_blocked_output_error()); + } + Err(error) => break 'transport Err(error), + } } Ok(InitializeOutcome::Rejected) => {} Err(e) => { @@ -268,17 +331,36 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { continue; } - if let Some(session_id) = session_id_from_message(&msg) - && state.open_session_streams.insert(session_id.clone()) - { - lifecycle.start_sse(Some(session_id), sse_event_tx.clone()); + if let Some(session_id) = session_id_from_message(&msg) { + for session_id in state.register_session_streams([session_id]) { + match lifecycle + .start_sse( + Some(session_id), + sse_event_tx.clone(), + SseStartContext { + events: &mut sse_event_rx, + outgoing: &mut outgoing, + buffered_outgoing: &mut buffered_outgoing, + posts: &mut posts, + state: &mut state, + }, + ) + .await + { + Ok(SseStartOutcome::Established) => {} + Ok(SseStartOutcome::OutgoingClosed) => { + break 'transport Err(sse_setup_blocked_output_error()); + } + Err(error) => break 'transport Err(error), + } + } } match state.prepare_post(msg) { // Responses answer SSE-delivered callbacks and must not be blocked // behind a POST that may be waiting for that callback response. - Ok(post) if is_response_only => response_posts.push(post), - Ok(post) => ordered_posts.push(post), + Ok(post) if is_response_only => posts.responses.push(post), + Ok(post) => posts.ordered.push(post), Err(e) => { error!("POST failed: {e}"); break Err(AcpError::internal_error().data(format!("POST: {e}"))); @@ -290,6 +372,56 @@ async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> { result } +fn sse_failure_error(failure: SseFailure) -> AcpError { + let scope = failure.session_id.as_deref().unwrap_or("connection"); + error!(session_id = ?failure.session_id, error = %failure.error, "SSE stream ended"); + AcpError::internal_error().data(format!("{scope} SSE stream ended: {}", failure.error)) +} + +fn sse_setup_blocked_output_error() -> AcpError { + AcpError::internal_error() + .data("outgoing channel closed while accepted messages awaited SSE stream establishment") +} + +fn handle_completed_post( + state: &mut ClientState, + completed: CompletedPost, +) -> Result<(), AcpError> { + let CompletedPost { + pending_requests, + result, + } = completed; + if let Err(error) = result { + state.remove_pending_requests(&pending_requests); + error!("POST failed: {error}"); + Err(AcpError::internal_error().data(format!("POST: {error}"))) + } else { + Ok(()) + } +} + +fn queue_response_post( + state: &mut ClientState, + posts: &mut PostQueues, + frame: TransportFrame, +) -> Result<(), AcpError> { + let post = match frame { + TransportFrame::Single(message) => state.prepare_post(message), + frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => { + state.prepare_frame_post(frame).map(|(post, session_ids)| { + debug_assert!(session_ids.is_empty()); + post + }) + } + } + .map_err(|error| { + error!("POST failed: {error}"); + AcpError::internal_error().data(format!("POST: {error}")) + })?; + posts.responses.push(post); + Ok(()) +} + fn is_response_only_frame(frame: &TransportFrame) -> bool { match frame { TransportFrame::Single(RawJsonRpcMessage::Response(_)) => true, @@ -417,6 +549,20 @@ struct HttpTransportLifecycle { sse_tasks: SseTasks, } +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +enum SseStartOutcome { + Established, + OutgoingClosed, +} + +struct SseStartContext<'a> { + events: &'a mut mpsc::UnboundedReceiver, + outgoing: &'a mut mpsc::UnboundedReceiver, + buffered_outgoing: &'a mut VecDeque, + posts: &'a mut PostQueues, + state: &'a mut ClientState, +} + impl HttpTransportLifecycle { fn new(connection: HttpConnection) -> Self { Self { @@ -425,9 +571,92 @@ impl HttpTransportLifecycle { } } - fn start_sse(&mut self, session_id: Option, event_tx: UnboundedSender) { - self.sse_tasks - .push(run_sse(self.connection.clone(), session_id, event_tx)); + async fn start_sse( + &mut self, + session_id: Option, + event_tx: UnboundedSender, + context: SseStartContext<'_>, + ) -> Result { + let SseStartContext { + events, + outgoing, + buffered_outgoing, + posts, + state, + } = context; + let mut establishing = FuturesUnordered::new(); + establishing.push(self.begin_sse(session_id, event_tx.clone())); + + loop { + if establishing.is_empty() { + return Ok(SseStartOutcome::Established); + } + let outcome = { + let failure = self.sse_tasks.next_failure().fuse(); + let established_next = establishing.next().fuse(); + let sse_event_next = events.next().fuse(); + let outgoing_next = outgoing.next().fuse(); + let ordered_post_next = posts.ordered.next_completion().fuse(); + let response_post_next = posts.responses.next_completion().fuse(); + pin_mut!( + failure, + established_next, + sse_event_next, + outgoing_next, + ordered_post_next, + response_post_next + ); + futures::select_biased! { + failure = failure => SseStartWait::Failure(failure), + established = established_next => SseStartWait::Established(established), + event = sse_event_next => SseStartWait::SseEvent(event), + post = response_post_next => SseStartWait::Post(post), + post = ordered_post_next => SseStartWait::Post(post), + outgoing = outgoing_next => SseStartWait::Outgoing(outgoing), + } + }; + match outcome { + SseStartWait::Established(Some(Ok(()))) => {} + SseStartWait::Established(Some(Err(_))) => { + return Err(sse_failure_error(self.sse_tasks.next_failure().await)); + } + SseStartWait::Established(None) => { + return Ok(SseStartOutcome::Established); + } + SseStartWait::Failure(failure) => return Err(sse_failure_error(failure)), + SseStartWait::SseEvent(Some(event)) => { + let open_session_ids = state.sessions_to_open_for_responses(&event.frame); + state.deliver_frame(event.frame); + for session_id in open_session_ids { + establishing.push(self.begin_sse(Some(session_id), event_tx.clone())); + } + } + SseStartWait::SseEvent(None) => { + return Err(AcpError::internal_error().data("SSE event channel closed")); + } + SseStartWait::Post(completed) => handle_completed_post(state, completed)?, + SseStartWait::Outgoing(Some(frame)) if is_response_only_frame(&frame) => { + queue_response_post(state, posts, frame)?; + } + SseStartWait::Outgoing(Some(frame)) => buffered_outgoing.push_back(frame), + SseStartWait::Outgoing(None) => return Ok(SseStartOutcome::OutgoingClosed), + } + } + } + + fn begin_sse( + &mut self, + session_id: Option, + event_tx: UnboundedSender, + ) -> futures::channel::oneshot::Receiver<()> { + let (established_tx, established_rx) = futures::channel::oneshot::channel(); + self.sse_tasks.push(run_sse( + self.connection.clone(), + session_id, + event_tx, + established_tx, + )); + established_rx } async fn next_sse_failure(&mut self) -> SseFailure { @@ -440,6 +669,14 @@ impl HttpTransportLifecycle { } } +enum SseStartWait { + Established(Option>), + Failure(SseFailure), + SseEvent(Option), + Post(CompletedPost), + Outgoing(Option), +} + impl Drop for HttpTransportLifecycle { fn drop(&mut self) { self.sse_tasks.abort_all(); @@ -451,10 +688,11 @@ fn run_sse( connection: HttpConnection, session_id: Option, event_tx: UnboundedSender, + established_tx: futures::channel::oneshot::Sender<()>, ) -> BoxFuture<'static, SseFailure> { Box::pin(async move { let label = session_id.clone(); - let error = match read_sse(connection, session_id, event_tx).await { + let error = match read_sse(connection, session_id, event_tx, established_tx).await { Ok(()) => "SSE stream closed".to_string(), Err(e) => e, }; @@ -493,24 +731,24 @@ impl SseTasks { struct ClientState { connection: HttpConnection, open_session_streams: HashSet, - pending_requests: HashMap, + pending_requests: HashMap>, incoming: futures::channel::mpsc::UnboundedSender, } struct PendingPost { - pending_request: Option<(RequestId, String)>, + pending_requests: Vec<(RequestId, String)>, response: BoxFuture<'static, Result<(), String>>, } impl PendingPost { fn into_completion(self) -> BoxFuture<'static, CompletedPost> { let Self { - pending_request, + pending_requests, response, } = self; async move { CompletedPost { - pending_request, + pending_requests, result: response.await, } } @@ -520,7 +758,7 @@ impl PendingPost { #[derive(Debug)] struct CompletedPost { - pending_request: Option<(RequestId, String)>, + pending_requests: Vec<(RequestId, String)>, result: Result<(), String>, } @@ -530,6 +768,18 @@ struct PostQueue { in_flight: Option>, } +#[derive(Default)] +struct PostQueues { + ordered: PostQueue, + responses: PostQueue, +} + +impl PostQueues { + fn is_empty(&self) -> bool { + self.ordered.is_empty() && self.responses.is_empty() + } +} + impl PostQueue { fn push(&mut self, post: PendingPost) { self.queued.push_back(post); @@ -621,16 +871,7 @@ impl ClientState { } fn prepare_post(&mut self, msg: RawJsonRpcMessage) -> Result { - let session_id = match method_for_message(&msg) { - Some(method) => { - let session_id = session_id_from_message(&msg); - if method_requires_session_header(method) && session_id.is_none() { - return Err(format!("method `{method}` requires sessionId in params")); - } - session_id - } - None => None, - }; + let session_id = validated_session_id(&msg)?; let connection_id = self .connection .connection_id() @@ -645,10 +886,10 @@ impl ClientState { request = request.header(HEADER_SESSION_ID, session_id); } - let pending_request = pending_request_for_message(&msg); - if let Some((id, method)) = &pending_request { - self.pending_requests.insert(id.clone(), method.clone()); - } + let pending_requests = pending_request_for_message(&msg) + .into_iter() + .collect::>(); + self.track_pending_requests(&pending_requests); let response = async move { let response = request.send().await.map_err(|e| e.to_string())?; @@ -660,12 +901,16 @@ impl ClientState { Ok(()) }; Ok(PendingPost { - pending_request, + pending_requests, response: response.boxed(), }) } - fn prepare_frame_post(&self, frame: TransportFrame) -> Result { + fn prepare_frame_post( + &mut self, + frame: TransportFrame, + ) -> Result<(PendingPost, Vec), String> { + let bookkeeping = FrameBookkeeping::for_frame(&frame)?; let connection_id = self .connection .connection_id() @@ -687,16 +932,78 @@ impl ClientState { } Ok(()) }; - Ok(PendingPost { - pending_request: None, - response: response.boxed(), - }) + self.track_pending_requests(&bookkeeping.pending_requests); + let session_ids = self.register_session_streams(bookkeeping.session_ids); + Ok(( + PendingPost { + pending_requests: bookkeeping.pending_requests, + response: response.boxed(), + }, + session_ids, + )) } - fn remove_pending_request(&mut self, pending_request: Option<&(RequestId, String)>) { - if let Some((id, _)) = pending_request { + fn track_pending_requests(&mut self, pending_requests: &[(RequestId, String)]) { + for (id, method) in pending_requests { + self.pending_requests + .entry(id.clone()) + .or_default() + .push_back(method.clone()); + } + } + + fn remove_pending_requests(&mut self, pending_requests: &[(RequestId, String)]) { + for (id, method) in pending_requests.iter().rev() { + let remove_entry = self.pending_requests.get_mut(id).is_some_and(|methods| { + if let Some(index) = methods.iter().rposition(|candidate| candidate == method) { + methods.remove(index); + } + methods.is_empty() + }); + if remove_entry { + self.pending_requests.remove(id); + } + } + } + + fn take_pending_request_method(&mut self, id: &RequestId) -> Option { + let (method, remove_entry) = { + let methods = self.pending_requests.get_mut(id)?; + (methods.pop_front(), methods.is_empty()) + }; + if remove_entry { self.pending_requests.remove(id); } + method + } + + fn register_session_streams( + &mut self, + session_ids: impl IntoIterator, + ) -> Vec { + session_ids + .into_iter() + .filter(|session_id| self.open_session_streams.insert(session_id.clone())) + .collect() + } + + fn sessions_to_open_for_responses(&mut self, frame: &TransportFrame) -> Vec { + match frame { + TransportFrame::Single(message) => self + .session_to_open_for_response(message) + .into_iter() + .collect(), + TransportFrame::Batch(batch) => batch + .entries() + .filter_map(|entry| match entry { + TransportBatchEntry::Message(message) => { + self.session_to_open_for_response(message) + } + TransportBatchEntry::Malformed { .. } => None, + }) + .collect(), + TransportFrame::Malformed { .. } => Vec::new(), + } } fn session_to_open_for_response(&mut self, msg: &RawJsonRpcMessage) -> Option { @@ -704,7 +1011,7 @@ impl ClientState { return None; }; let id = msg.response_id().and_then(pending_request_key)?; - let method = self.pending_requests.remove(&id); + let method = self.take_pending_request_method(&id); if !method.as_deref().is_some_and(is_session_opening_method) { return None; @@ -735,6 +1042,53 @@ impl ClientState { } } +#[derive(Default)] +struct FrameBookkeeping { + session_ids: Vec, + pending_requests: Vec<(RequestId, String)>, +} + +impl FrameBookkeeping { + fn for_frame(frame: &TransportFrame) -> Result { + let mut bookkeeping = Self::default(); + match frame { + TransportFrame::Single(message) => bookkeeping.add_message(message)?, + TransportFrame::Batch(batch) => { + for entry in batch.entries() { + if let TransportBatchEntry::Message(message) = entry { + bookkeeping.add_message(message)?; + } + } + } + TransportFrame::Malformed { .. } => {} + } + Ok(bookkeeping) + } + + fn add_message(&mut self, message: &RawJsonRpcMessage) -> Result<(), String> { + if let Some(session_id) = validated_session_id(message)? + && !self.session_ids.contains(&session_id) + { + self.session_ids.push(session_id); + } + if let Some(pending_request) = pending_request_for_message(message) { + self.pending_requests.push(pending_request); + } + Ok(()) + } +} + +fn validated_session_id(msg: &RawJsonRpcMessage) -> Result, String> { + let Some(method) = method_for_message(msg) else { + return Ok(None); + }; + let session_id = session_id_from_message(msg); + if method_requires_session_header(method) && session_id.is_none() { + return Err(format!("method `{method}` requires sessionId in params")); + } + Ok(session_id) +} + fn is_session_opening_method(method: &str) -> bool { matches!(method, "session/new" | "session/fork") } @@ -743,6 +1097,7 @@ async fn read_sse( connection: HttpConnection, session_id: Option, event_tx: UnboundedSender, + established_tx: futures::channel::oneshot::Sender<()>, ) -> Result<(), String> { let connection_id = connection .connection_id() @@ -760,6 +1115,7 @@ async fn read_sse( return Err(format!("HTTP {}", response.status())); } trace!(session_id = ?session_id, "SSE stream open"); + let _ = established_tx.send(()); let mut events = eventsource_stream::EventStream::new(response.bytes_stream()); while let Some(event) = events.next().await { @@ -906,7 +1262,7 @@ mod tests { convert::Infallible, sync::{ Arc, - atomic::{AtomicUsize, Ordering}, + atomic::{AtomicBool, AtomicUsize, Ordering}, }, time::Duration, }; @@ -936,6 +1292,11 @@ mod tests { >, } + struct InitializeThenExitClient { + sse_started: Arc, + finished: Arc, + } + struct QueueOutgoingThenText { text: Option, outgoing: Option>, @@ -1007,6 +1368,110 @@ mod tests { assert!(!is_response_only_frame(&call_shaped)); } + fn initialized_client_state() -> ClientState { + let connection = HttpConnection::new( + url::Url::parse("http://127.0.0.1/acp").unwrap(), + reqwest::Client::new(), + ); + connection.set_connection_id("connection-1".to_string()); + let (incoming, _incoming_rx) = mpsc::unbounded(); + ClientState { + connection, + open_session_streams: HashSet::new(), + pending_requests: HashMap::new(), + incoming, + } + } + + #[test] + fn batch_post_validation_happens_before_tracking_requests_or_sessions() { + let mut state = initialized_client_state(); + let frame = TransportFrame::Batch( + TransportBatch::from_messages([ + RawJsonRpcMessage::request( + "custom/valid".to_string(), + json!({}), + RequestId::Number(1), + ) + .unwrap(), + RawJsonRpcMessage::request( + "session/prompt".to_string(), + json!({ "prompt": [] }), + RequestId::Number(2), + ) + .unwrap(), + ]) + .unwrap(), + ); + + let Err(error) = state.prepare_frame_post(frame) else { + panic!("batch should require sessionId for session/prompt"); + }; + + assert_eq!( + error, + "method `session/prompt` requires sessionId in params" + ); + assert!(state.pending_requests.is_empty()); + assert!(state.open_session_streams.is_empty()); + } + + #[test] + fn batch_post_tracks_every_non_null_request_and_rolls_back_from_the_back() { + let mut state = initialized_client_state(); + state.track_pending_requests(&[(RequestId::Number(7), "session/fork".to_string())]); + let frame = TransportFrame::Batch( + TransportBatch::from_messages([ + RawJsonRpcMessage::request( + "session/fork".to_string(), + json!({ "sessionId": "source-a" }), + RequestId::Number(7), + ) + .unwrap(), + RawJsonRpcMessage::request( + "custom/request".to_string(), + json!({ "sessionId": "source-b" }), + RequestId::Number(7), + ) + .unwrap(), + RawJsonRpcMessage::request( + "session/fork".to_string(), + json!({ "sessionId": "source-a" }), + RequestId::Null, + ) + .unwrap(), + ]) + .unwrap(), + ); + + let (post, session_ids) = state.prepare_frame_post(frame).unwrap(); + + assert_eq!(session_ids, ["source-a", "source-b"]); + assert_eq!( + state.pending_requests.get(&RequestId::Number(7)).unwrap(), + &VecDeque::from([ + "session/fork".to_string(), + "session/fork".to_string(), + "custom/request".to_string(), + ]) + ); + assert_eq!( + post.pending_requests, + [ + (RequestId::Number(7), "session/fork".to_string()), + (RequestId::Number(7), "custom/request".to_string()), + ] + ); + assert!(!state.pending_requests.contains_key(&RequestId::Null)); + + state.remove_pending_requests(&post.pending_requests); + + assert_eq!( + state.pending_requests.get(&RequestId::Number(7)).unwrap(), + &VecDeque::from(["session/fork".to_string()]) + ); + } + impl WsSink for RecordingWsSink { async fn send(&mut self, message: WsMessage) -> Result<(), String> { self.0 @@ -1125,6 +1590,41 @@ mod tests { } } + impl ConnectTo for InitializeThenExitClient { + async fn connect_to(self, agent: impl ConnectTo) -> Result<(), AcpError> { + let Self { + sse_started, + finished, + } = self; + let (mut channel, transport) = agent.into_channel_and_future(); + let client = async move { + channel + .tx + .unbounded_send(single_frame( + RawJsonRpcMessage::request( + "initialize".to_string(), + json!({}), + RequestId::Number(1), + ) + .unwrap(), + )) + .map_err(|error| { + AcpError::internal_error().data(format!("send initialize: {error}")) + })?; + into_single_message(channel.rx.next().await.ok_or_else(|| { + AcpError::internal_error().data("initialize response channel closed") + })?)?; + + sse_started.notified().await; + finished.notify_one(); + Ok(()) + }; + + let ((), ()) = futures::try_join!(transport, client)?; + Ok(()) + } + } + #[test] fn new_targets_standard_acp_endpoint() { assert_eq!( @@ -1395,6 +1895,172 @@ mod tests { server.abort(); } + #[tokio::test] + async fn batch_fork_opens_source_and_result_session_streams() { + let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel(); + let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel(); + let post_count = Arc::new(AtomicUsize::new(0)); + let emit_response = Arc::new(Notify::new()); + let connection_stream_established = Arc::new(AtomicBool::new(false)); + let source_stream_established = Arc::new(AtomicBool::new(false)); + let response_batch = json!([ + { + "jsonrpc": "2.0", + "id": 2, + "result": { "sessionId": "forked-session" } + } + ]); + let app = Router::new().route( + "/acp", + post({ + let post_count = post_count.clone(); + let connection_stream_established = connection_stream_established.clone(); + let source_stream_established = source_stream_established.clone(); + move |body: String| { + let post_count = post_count.clone(); + let post_tx = post_tx.clone(); + let connection_stream_established = connection_stream_established.clone(); + let source_stream_established = source_stream_established.clone(); + async move { + if post_count.fetch_add(1, Ordering::SeqCst) == 0 { + return initialize_response().await.into_response(); + } + + if !connection_stream_established.load(Ordering::SeqCst) + || !source_stream_established.load(Ordering::SeqCst) + { + return StatusCode::CONFLICT.into_response(); + } + post_tx + .send(serde_json::from_str::(&body).unwrap()) + .unwrap(); + StatusCode::ACCEPTED.into_response() + } + } + }) + .get({ + let emit_response = emit_response.clone(); + let response_batch = response_batch.clone(); + let connection_stream_established = connection_stream_established.clone(); + let source_stream_established = source_stream_established.clone(); + move |headers: HeaderMap| { + let emit_response = emit_response.clone(); + let response_batch = response_batch.clone(); + let get_tx = get_tx.clone(); + let connection_stream_established = connection_stream_established.clone(); + let source_stream_established = source_stream_established.clone(); + async move { + let session_id = headers + .get(HEADER_SESSION_ID) + .and_then(|value| value.to_str().ok()) + .map(String::from); + let is_connection_stream = session_id.is_none(); + let is_source_stream = session_id.as_deref() == Some("source-session"); + if is_connection_stream { + sleep(Duration::from_millis(50)).await; + connection_stream_established.store(true, Ordering::SeqCst); + } + if is_source_stream { + sleep(Duration::from_millis(50)).await; + source_stream_established.store(true, Ordering::SeqCst); + } + get_tx.send(session_id).unwrap(); + + let stream = async_stream::stream! { + if is_source_stream { + emit_response.notified().await; + yield Ok::<_, Infallible>( + Event::default().data(response_batch.to_string()), + ); + } + futures::future::pending::<()>().await; + }; + Sse::new(stream) + } + } + }) + .delete(|| async { StatusCode::ACCEPTED }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + let client = HttpClient::new(format!("http://{addr}")).unwrap(); + let (mut caller, transport) = Channel::duplex(); + let transport = tokio::spawn(run(client, transport)); + + caller + .tx + .unbounded_send(single_frame( + RawJsonRpcMessage::request( + "initialize".to_string(), + json!({}), + RequestId::Number(1), + ) + .unwrap(), + )) + .unwrap(); + timeout(Duration::from_secs(1), caller.rx.next()) + .await + .unwrap() + .unwrap(); + + caller + .tx + .unbounded_send(TransportFrame::Batch( + TransportBatch::from_messages([RawJsonRpcMessage::request( + "session/fork".to_string(), + json!({ "sessionId": "source-session" }), + RequestId::Number(2), + ) + .unwrap()]) + .unwrap(), + )) + .unwrap(); + + let connection_stream = timeout(Duration::from_secs(1), get_rx.recv()) + .await + .unwrap() + .unwrap(); + assert!(connection_stream.is_none()); + let source_stream = timeout(Duration::from_secs(1), get_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(source_stream.as_deref(), Some("source-session")); + let posted = timeout(Duration::from_secs(1), post_rx.recv()) + .await + .unwrap() + .unwrap(); + assert!(posted.is_array(), "outgoing batch must remain an array"); + + emit_response.notify_one(); + let response = timeout(Duration::from_secs(1), caller.rx.next()) + .await + .unwrap() + .unwrap(); + assert!(matches!(&response, TransportFrame::Batch(_))); + assert_eq!( + serde_json::from_str::(&response.to_json().unwrap()).unwrap(), + response_batch + ); + let forked_stream = timeout(Duration::from_secs(1), get_rx.recv()) + .await + .unwrap() + .unwrap(); + assert_eq!(forked_stream.as_deref(), Some("forked-session")); + + drop(caller); + timeout(Duration::from_secs(1), transport) + .await + .unwrap() + .unwrap() + .unwrap(); + + server.abort(); + } + #[tokio::test] async fn custom_response_with_session_id_does_not_open_session_sse() { let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel(); @@ -1880,6 +2546,312 @@ mod tests { server.abort(); } + #[tokio::test] + async fn client_completion_cancels_pending_sse_establishment() { + let sse_started = Arc::new(Notify::new()); + let delete_count = Arc::new(AtomicUsize::new(0)); + let client_finished = Arc::new(Notify::new()); + let app = Router::new().route( + "/acp", + post(initialize_response) + .get({ + let sse_started = sse_started.clone(); + move || { + let sse_started = sse_started.clone(); + async move { + sse_started.notify_one(); + futures::future::pending::().await + } + } + }) + .delete({ + let delete_count = delete_count.clone(); + move || { + let delete_count = delete_count.clone(); + async move { + delete_count.fetch_add(1, Ordering::SeqCst); + StatusCode::ACCEPTED + } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + let client = HttpClient::new(format!("http://{addr}")).unwrap(); + let connection = tokio::spawn(client.connect_to(InitializeThenExitClient { + sse_started, + finished: client_finished.clone(), + })); + + timeout(Duration::from_secs(1), client_finished.notified()) + .await + .expect("client foreground did not finish after the SSE request started"); + + timeout(Duration::from_secs(1), connection) + .await + .expect("transport remained blocked on SSE response headers") + .unwrap() + .unwrap(); + assert_eq!(delete_count.load(Ordering::SeqCst), 1); + + server.abort(); + } + + #[tokio::test] + async fn stalled_sse_establishment_observes_earlier_post_failure() { + let app = Router::new().route( + "/acp", + get(|| async { futures::future::pending::().await }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let connection = HttpConnection::new( + url::Url::parse(&format!("http://{addr}/acp")).unwrap(), + reqwest::Client::new(), + ); + connection.set_connection_id("connection-1".to_string()); + let (incoming, _incoming_rx) = mpsc::unbounded(); + let mut state = ClientState { + connection: connection.clone(), + open_session_streams: HashSet::new(), + pending_requests: HashMap::new(), + incoming, + }; + let pending_request = (RequestId::Number(7), "custom/earlier".to_string()); + state.track_pending_requests(std::slice::from_ref(&pending_request)); + let mut posts = PostQueues::default(); + posts.ordered.push(PendingPost { + pending_requests: vec![pending_request], + response: async { Err("earlier post failed".to_string()) }.boxed(), + }); + + let (_outgoing_tx, mut outgoing) = mpsc::unbounded(); + let mut buffered_outgoing = VecDeque::new(); + let (event_tx, mut event_rx) = mpsc::unbounded(); + let mut lifecycle = HttpTransportLifecycle::new(connection); + let error = timeout( + Duration::from_secs(1), + lifecycle.start_sse( + Some("later-session".to_string()), + event_tx, + SseStartContext { + events: &mut event_rx, + outgoing: &mut outgoing, + buffered_outgoing: &mut buffered_outgoing, + posts: &mut posts, + state: &mut state, + }, + ), + ) + .await + .expect("stalled SSE setup hid an earlier POST failure") + .unwrap_err(); + + assert!(error.to_string().contains("earlier post failed")); + assert!(state.pending_requests.is_empty()); + + lifecycle.close().await; + server.abort(); + } + + #[tokio::test] + async fn stalled_sse_establishment_keeps_callback_responses_moving() { + let release_get = Arc::new(Notify::new()); + let complete_earlier_post = Arc::new(Notify::new()); + let app = Router::new().route( + "/acp", + post({ + let release_get = release_get.clone(); + let complete_earlier_post = complete_earlier_post.clone(); + move || { + release_get.notify_one(); + complete_earlier_post.notify_one(); + async { StatusCode::ACCEPTED } + } + }) + .get({ + let release_get = release_get.clone(); + move || { + let release_get = release_get.clone(); + async move { + release_get.notified().await; + Sse::new(futures::stream::pending::>()) + } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + + let connection = HttpConnection::new( + url::Url::parse(&format!("http://{addr}/acp")).unwrap(), + reqwest::Client::new(), + ); + connection.set_connection_id("connection-1".to_string()); + let (incoming, mut incoming_rx) = mpsc::unbounded(); + let mut state = ClientState { + connection: connection.clone(), + open_session_streams: HashSet::new(), + pending_requests: HashMap::new(), + incoming, + }; + let mut posts = PostQueues::default(); + posts.ordered.push(PendingPost { + pending_requests: Vec::new(), + response: async move { + complete_earlier_post.notified().await; + Ok(()) + } + .boxed(), + }); + + let (outgoing_tx, mut outgoing) = mpsc::unbounded(); + let outgoing_guard = outgoing_tx.clone(); + let mut buffered_outgoing = VecDeque::new(); + let (event_tx, mut event_rx) = mpsc::unbounded(); + event_tx + .unbounded_send(SseMessage { + frame: single_frame( + RawJsonRpcMessage::request( + "test/callback".to_string(), + json!({}), + RequestId::Number(99), + ) + .unwrap(), + ), + }) + .unwrap(); + + let responder = async move { + let callback = incoming_rx + .next() + .await + .expect("callback was not delivered"); + assert!(matches!( + into_single_message(callback).unwrap(), + RawJsonRpcMessage::Request(request) + if request.method.as_ref() == "test/callback" + )); + outgoing_tx + .unbounded_send(single_frame(RawJsonRpcMessage::response( + RequestId::Number(99), + Ok(json!({})), + ))) + .unwrap(); + }; + let mut lifecycle = HttpTransportLifecycle::new(connection); + let (outcome, ()) = timeout(Duration::from_secs(1), async { + futures::join!( + lifecycle.start_sse( + Some("later-session".to_string()), + event_tx, + SseStartContext { + events: &mut event_rx, + outgoing: &mut outgoing, + buffered_outgoing: &mut buffered_outgoing, + posts: &mut posts, + state: &mut state, + }, + ), + responder, + ) + }) + .await + .expect("callback response deadlocked behind stalled SSE establishment"); + + assert_eq!(outcome.unwrap(), SseStartOutcome::Established); + assert!(buffered_outgoing.is_empty()); + + drop(outgoing_guard); + lifecycle.close().await; + server.abort(); + } + + #[tokio::test] + async fn pending_sse_establishment_reports_buffered_output_on_shutdown() { + let sse_started = Arc::new(Notify::new()); + let delete_count = Arc::new(AtomicUsize::new(0)); + let app = Router::new().route( + "/acp", + post(initialize_response) + .get({ + let sse_started = sse_started.clone(); + move || { + let sse_started = sse_started.clone(); + async move { + sse_started.notify_one(); + futures::future::pending::().await + } + } + }) + .delete({ + let delete_count = delete_count.clone(); + move || { + let delete_count = delete_count.clone(); + async move { + delete_count.fetch_add(1, Ordering::SeqCst); + StatusCode::ACCEPTED + } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + let client = HttpClient::new(format!("http://{addr}")).unwrap(); + let (mut caller, transport) = Channel::duplex(); + let transport = tokio::spawn(run(client, transport)); + + caller + .tx + .unbounded_send(single_frame( + RawJsonRpcMessage::request( + "initialize".to_string(), + json!({}), + RequestId::Number(1), + ) + .unwrap(), + )) + .unwrap(); + timeout(Duration::from_secs(1), caller.rx.next()) + .await + .unwrap() + .unwrap(); + timeout(Duration::from_secs(1), sse_started.notified()) + .await + .expect("connection SSE request did not reach the server"); + + caller + .tx + .unbounded_send(single_frame( + RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(), + )) + .unwrap(); + drop(caller); + + let error = timeout(Duration::from_secs(1), transport) + .await + .expect("transport remained blocked on SSE response headers") + .unwrap() + .unwrap_err(); + assert!(error.to_string().contains("accepted messages")); + assert_eq!(delete_count.load(Ordering::SeqCst), 1); + + server.abort(); + } + #[tokio::test] async fn sse_continues_while_post_is_pending() { let post_started = Arc::new(Notify::new()); From c75a71e4e63b141f23c0d1754d754eb875d1b18f Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Thu, 23 Jul 2026 15:08:11 +0200 Subject: [PATCH 2/3] fix(http): drain routed frames on shutdown --- .../src/connection.rs | 95 ++++- .../src/http_server.rs | 75 ++++ .../src/websocket_server.rs | 370 ++++++++++++++++-- 3 files changed, 513 insertions(+), 27 deletions(-) diff --git a/src/agent-client-protocol-http/src/connection.rs b/src/agent-client-protocol-http/src/connection.rs index de9292d..50b7488 100644 --- a/src/agent-client-protocol-http/src/connection.rs +++ b/src/agent-client-protocol-http/src/connection.rs @@ -555,10 +555,13 @@ async fn close_connection_task(connection: Weak) { let Some(connection) = connection.upgrade() else { return; }; - connection.close_streams(); - if let Some(h) = connection.router_handle.lock().await.take() { - h.abort(); + let router_handle = connection.router_handle.lock().await.take(); + if let Some(h) = router_handle + && let Err(error) = h.await + { + error!("HTTP outbound router task failed while draining: {error}"); } + connection.close_streams(); } fn pending_route_key(id: &RequestId) -> Option { @@ -739,6 +742,38 @@ mod tests { } } + struct FinalFrameThenExitAgentFactory { + emit: Arc, + } + + impl AgentFactory for FinalFrameThenExitAgentFactory { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (agent, transport) = Channel::duplex(); + let emit = self.emit.clone(); + let future = Box::pin(async move { + emit.notified().await; + agent + .tx + .unbounded_send(TransportFrame::Single( + RawJsonRpcMessage::notification( + "test/final".to_string(), + serde_json::json!({}), + ) + .unwrap(), + )) + .unwrap(); + Ok(()) + }); + + (transport, future) + } + } + #[tokio::test] async fn agent_exit_removes_connection_and_closes_streams() { let exit = Arc::new(Notify::new()); @@ -821,6 +856,60 @@ mod tests { assert!(*connection.subscribe_closed().borrow()); } + #[tokio::test] + async fn agent_exit_waits_for_router_to_flush_final_frame() { + let emit = Arc::new(Notify::new()); + let registry = ConnectionRegistry::new(Arc::new(FinalFrameThenExitAgentFactory { + emit: emit.clone(), + })); + let (connection_id, connection) = registry.create_connection().await; + let (_replay, mut outbound) = connection.subscribe_connection_stream().await; + connection.start_router().await; + + let stream = match &connection.outbound_transport { + OutboundTransport::Http(http) => http.connection_stream.clone(), + OutboundTransport::WebSocket(_) => unreachable!("created an HTTP connection"), + }; + let state_guard = stream.state.lock().await; + emit.notify_one(); + + timeout(Duration::from_secs(1), async { + loop { + if registry.get(&connection_id).await.is_none() { + break; + } + sleep(Duration::from_millis(10)).await; + } + }) + .await + .unwrap(); + assert!( + !*connection.subscribe_closed().borrow(), + "stream closure must wait for the blocked outbound router" + ); + + drop(state_guard); + let text = timeout(Duration::from_secs(1), outbound.recv()) + .await + .unwrap() + .expect("final frame should reach the established stream"); + let message = serde_json::from_str::(&text).unwrap(); + assert!(matches!( + message, + RawJsonRpcMessage::Notification(notification) + if notification.method.as_ref() == "test/final" + )); + + timeout(Duration::from_secs(1), async { + let mut closed = connection.subscribe_closed(); + while !*closed.borrow() { + closed.changed().await.unwrap(); + } + }) + .await + .unwrap(); + } + #[tokio::test] async fn protocol_level_notification_routes_to_connection_stream() { let exit = Arc::new(Notify::new()); diff --git a/src/agent-client-protocol-http/src/http_server.rs b/src/agent-client-protocol-http/src/http_server.rs index 2d47893..db85539 100644 --- a/src/agent-client-protocol-http/src/http_server.rs +++ b/src/agent-client-protocol-http/src/http_server.rs @@ -357,10 +357,19 @@ pub(crate) async fn handle_get( yield Ok::<_, Infallible>(Event::default().data(msg)); } loop { + while let Ok(msg) = receiver.try_recv() { + trace!(payload = %msg, "SSE → client"); + yield Ok(Event::default().data(msg)); + } if *closed.borrow() { + while let Ok(msg) = receiver.try_recv() { + trace!(payload = %msg, "SSE → client"); + yield Ok(Event::default().data(msg)); + } break; } tokio::select! { + biased; recv = receiver.recv() => match recv { Some(msg) => { trace!(payload = %msg, "SSE → client"); @@ -370,6 +379,10 @@ pub(crate) async fn handle_get( }, changed = closed.changed() => { if changed.is_err() || *closed.borrow() { + while let Ok(msg) = receiver.try_recv() { + trace!(payload = %msg, "SSE → client"); + yield Ok(Event::default().data(msg)); + } break; } } @@ -1269,6 +1282,68 @@ mod tests { connection.shutdown().await; } + #[tokio::test] + async fn sse_drains_queued_message_before_connection_close() { + let (forwarded_tx, _forwarded_rx) = mpsc::unbounded_channel(); + let registry = Arc::new(ConnectionRegistry::new(Arc::new(CapturingAgentFactory { + forwarded: forwarded_tx, + }))); + let (connection_id, connection) = registry.create_connection().await; + let request = Request::builder() + .method("GET") + .uri("/acp") + .header(header::ACCEPT, EVENT_STREAM_MIME_TYPE) + .header(HEADER_CONNECTION_ID, connection_id) + .body(Body::empty()) + .unwrap(); + let response = handle_get(registry, request).await; + assert_eq!(response.status(), StatusCode::OK); + + connection + .push_connection_stream_for_test("final-message".to_string()) + .await; + connection.shutdown().await; + + let body = timeout( + Duration::from_secs(1), + axum::body::to_bytes(response.into_body(), 1024), + ) + .await + .expect("SSE body should close without hanging") + .unwrap(); + let body = String::from_utf8(body.to_vec()).unwrap(); + assert!(body.contains("data: final-message")); + } + + #[tokio::test] + async fn sse_subscribed_after_connection_close_ends_without_hanging() { + let (forwarded_tx, _forwarded_rx) = mpsc::unbounded_channel(); + let registry = Arc::new(ConnectionRegistry::new(Arc::new(CapturingAgentFactory { + forwarded: forwarded_tx, + }))); + let (connection_id, connection) = registry.create_connection().await; + connection.shutdown().await; + + let request = Request::builder() + .method("GET") + .uri("/acp") + .header(header::ACCEPT, EVENT_STREAM_MIME_TYPE) + .header(HEADER_CONNECTION_ID, connection_id) + .body(Body::empty()) + .unwrap(); + let response = handle_get(registry, request).await; + assert_eq!(response.status(), StatusCode::OK); + + let body = timeout( + Duration::from_secs(1), + axum::body::to_bytes(response.into_body(), 1024), + ) + .await + .expect("an already-closed SSE stream should end without hanging") + .unwrap(); + assert!(body.is_empty()); + } + #[tokio::test] async fn post_forwards_header_session_id_to_agent_params() { let (forwarded_tx, mut forwarded_rx) = mpsc::unbounded_channel(); diff --git a/src/agent-client-protocol-http/src/websocket_server.rs b/src/agent-client-protocol-http/src/websocket_server.rs index 8d18f2b..17f7074 100644 --- a/src/agent-client-protocol-http/src/websocket_server.rs +++ b/src/agent-client-protocol-http/src/websocket_server.rs @@ -65,24 +65,67 @@ async fn run_ws( } } + run_ws_message_loop( + &mut ws_tx, + &mut ws_rx, + &mut outbound_rx, + &mut closed, + &connection_id, + &connection, + ) + .await; + + debug!(connection_id = %connection_id, "Cleaning up WebSocket connection"); + if let Some(conn) = registry.remove(&connection_id).await { + conn.shutdown().await; + } +} + +async fn run_ws_message_loop( + ws_tx: &mut futures::stream::SplitSink, + ws_rx: &mut futures::stream::SplitStream, + outbound_rx: &mut tokio::sync::mpsc::Receiver, + closed: &mut tokio::sync::watch::Receiver, + connection_id: &str, + connection: &crate::connection::Connection, +) { loop { if *closed.borrow() { + drain_queued_outbound(ws_tx, outbound_rx, connection_id).await; break; } tokio::select! { + recv = outbound_rx.recv() => { + match recv { + Some(text) => { + if !send_outbound_text(ws_tx, text, connection_id).await { + break; + } + } + None => break, + } + } + + changed = closed.changed() => { + if changed.is_err() || *closed.borrow() { + drain_queued_outbound(ws_tx, outbound_rx, connection_id).await; + break; + } + } + msg_result = ws_rx.next() => { match msg_result { Some(Ok(WsMessage::Text(text))) => { - let text_str = text.to_string(); - trace!(connection_id = %connection_id, payload = %text_str, "Client → Agent: {} bytes", text_str.len()); - let frame = TransportFrame::parse_json(&text_str); - if let TransportFrame::Single(parsed) = &frame - && let Some(sid) = session_id_from_message(parsed) - && let RawJsonRpcMessage::Request(req) = parsed { - trace!(connection_id = %connection_id, session_id = %sid, request_id = ?req.id, "Client → Agent (session)"); - } - if connection.send_frame_to_agent(frame).is_err() { - error!(connection_id = %connection_id, "Agent channel closed"); + if !forward_client_text( + text.to_string(), + ws_tx, + outbound_rx, + closed, + connection_id, + connection, + ) + .await + { break; } } @@ -101,31 +144,96 @@ async fn run_ws( None => break, } } + } + } +} - recv = outbound_rx.recv() => { - match recv { - Some(text) => { - trace!(connection_id = %connection_id, payload = %text, "Agent → Client: {} bytes", text.len()); - if ws_tx.send(WsMessage::Text(text.into())).await.is_err() { - error!(connection_id = %connection_id, "WebSocket send failed"); - break; - } +async fn forward_client_text( + text: String, + ws_tx: &mut S, + outbound_rx: &mut tokio::sync::mpsc::Receiver, + closed: &mut tokio::sync::watch::Receiver, + connection_id: &str, + connection: &crate::connection::Connection, +) -> bool +where + S: futures::Sink + Unpin, +{ + trace!(connection_id = %connection_id, payload = %text, "Client → Agent: {} bytes", text.len()); + let frame = TransportFrame::parse_json(&text); + if let TransportFrame::Single(parsed) = &frame + && let Some(sid) = session_id_from_message(parsed) + && let RawJsonRpcMessage::Request(req) = parsed + { + trace!(connection_id = %connection_id, session_id = %sid, request_id = ?req.id, "Client → Agent (session)"); + } + if connection.send_frame_to_agent(frame).is_err() { + error!(connection_id = %connection_id, "Agent channel closed"); + drain_outbound_until_closed(ws_tx, outbound_rx, closed, connection_id).await; + false + } else { + true + } +} + +async fn drain_outbound_until_closed( + ws_tx: &mut S, + outbound_rx: &mut tokio::sync::mpsc::Receiver, + closed: &mut tokio::sync::watch::Receiver, + connection_id: &str, +) where + S: futures::Sink + Unpin, +{ + loop { + drain_queued_outbound(ws_tx, outbound_rx, connection_id).await; + if *closed.borrow() { + drain_queued_outbound(ws_tx, outbound_rx, connection_id).await; + break; + } + tokio::select! { + biased; + recv = outbound_rx.recv() => match recv { + Some(text) => { + if !send_outbound_text(ws_tx, text, connection_id).await { + break; } - None => break, } - } - + None => break, + }, changed = closed.changed() => { if changed.is_err() || *closed.borrow() { + drain_queued_outbound(ws_tx, outbound_rx, connection_id).await; break; } } } } +} - debug!(connection_id = %connection_id, "Cleaning up WebSocket connection"); - if let Some(conn) = registry.remove(&connection_id).await { - conn.shutdown().await; +async fn drain_queued_outbound( + ws_tx: &mut S, + outbound_rx: &mut tokio::sync::mpsc::Receiver, + connection_id: &str, +) where + S: futures::Sink + Unpin, +{ + while let Ok(text) = outbound_rx.try_recv() { + if !send_outbound_text(ws_tx, text, connection_id).await { + break; + } + } +} + +async fn send_outbound_text(ws_tx: &mut S, text: String, connection_id: &str) -> bool +where + S: futures::Sink + Unpin, +{ + trace!(connection_id = %connection_id, payload = %text, "Agent → Client: {} bytes", text.len()); + if ws_tx.send(WsMessage::Text(text.into())).await.is_err() { + error!(connection_id = %connection_id, "WebSocket send failed"); + false + } else { + true } } @@ -235,6 +343,220 @@ mod tests { } } + struct FinalFrameThenExitAgentFactory { + emit: Arc, + } + + impl AgentFactory for FinalFrameThenExitAgentFactory { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (agent, transport) = Channel::duplex(); + let emit = self.emit.clone(); + let future = Box::pin(async move { + emit.notified().await; + agent + .tx + .unbounded_send(TransportFrame::Single( + RawJsonRpcMessage::notification( + "test/final".to_string(), + serde_json::json!({}), + ) + .unwrap(), + )) + .unwrap(); + Ok(()) + }); + + (transport, future) + } + } + + struct FinalFrameAfterInputCloseAgentFactory { + emit: Arc, + } + + impl AgentFactory for FinalFrameAfterInputCloseAgentFactory { + fn spawn_agent( + &self, + ) -> ( + Channel, + BoxFuture<'static, agent_client_protocol::Result<()>>, + ) { + let (agent, transport) = Channel::duplex(); + let emit = self.emit.clone(); + let future = Box::pin(async move { + drop(agent.rx); + emit.notified().await; + agent + .tx + .unbounded_send(TransportFrame::Single( + RawJsonRpcMessage::notification( + "test/final".to_string(), + serde_json::json!({}), + ) + .unwrap(), + )) + .unwrap(); + Ok(()) + }); + + (transport, future) + } + } + + #[tokio::test] + async fn websocket_drains_final_agent_frame_before_closing() { + let emit = Arc::new(tokio::sync::Notify::new()); + let registry = Arc::new(ConnectionRegistry::new(Arc::new( + FinalFrameThenExitAgentFactory { emit: emit.clone() }, + ))); + let app = Router::new().route( + "/acp", + get({ + let registry = registry.clone(); + move |ws: WebSocketUpgrade| { + let registry = registry.clone(); + let emit = emit.clone(); + async move { + ws.on_upgrade(move |socket| async move { + let connection_id = ConnectionRegistry::next_connection_id(); + let connection = registry + .create_websocket_connection_with_id(connection_id.clone()) + .await; + connection.start_router().await; + let (replay, mut outbound_rx) = + connection.subscribe_all_outbound().await; + assert!(replay.is_empty()); + let mut closed = connection.subscribe_closed(); + + emit.notify_one(); + while !*closed.borrow() { + closed.changed().await.unwrap(); + } + + let (mut ws_tx, mut ws_rx) = socket.split(); + run_ws_message_loop( + &mut ws_tx, + &mut ws_rx, + &mut outbound_rx, + &mut closed, + &connection_id, + &connection, + ) + .await; + + if let Some(connection) = registry.remove(&connection_id).await { + connection.shutdown().await; + } + }) + } + } + }), + ); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + axum::serve(listener, app).await.unwrap(); + }); + let (mut client, _) = connect_async(format!("ws://{addr}/acp")).await.unwrap(); + + let frame = timeout(Duration::from_secs(1), client.next()) + .await + .unwrap() + .unwrap() + .unwrap(); + let ClientWsMessage::Text(text) = frame else { + panic!("expected final text frame: {frame:?}"); + }; + let message = serde_json::from_str::(&text).unwrap(); + assert!(matches!( + message, + RawJsonRpcMessage::Notification(notification) + if notification.method.as_ref() == "test/final" + )); + + server.abort(); + } + + #[tokio::test] + async fn inbound_after_agent_exit_drains_queued_final_frame() { + let emit = Arc::new(tokio::sync::Notify::new()); + let registry = ConnectionRegistry::new(Arc::new(FinalFrameAfterInputCloseAgentFactory { + emit: emit.clone(), + })); + let connection_id = ConnectionRegistry::next_connection_id(); + let connection = registry + .create_websocket_connection_with_id(connection_id.clone()) + .await; + connection.start_router().await; + let (replay, mut outbound_rx) = connection.subscribe_all_outbound().await; + assert!(replay.is_empty()); + let mut closed = connection.subscribe_closed(); + + timeout(Duration::from_secs(1), async { + loop { + let probe = RawJsonRpcMessage::notification( + "test/probe".to_string(), + serde_json::json!({}), + ) + .unwrap(); + if connection + .send_frame_to_agent(TransportFrame::Single(probe)) + .is_err() + { + break; + } + tokio::task::yield_now().await; + } + }) + .await + .expect("agent input did not close"); + assert!(!*closed.borrow(), "outbound routing should still be active"); + + let (mut ws_tx, mut ws_rx) = futures::channel::mpsc::unbounded::(); + let inbound = + RawJsonRpcMessage::notification("test/inbound".to_string(), serde_json::json!({})) + .unwrap(); + let forward = forward_client_text( + serde_json::to_string(&inbound).unwrap(), + &mut ws_tx, + &mut outbound_rx, + &mut closed, + &connection_id, + &connection, + ); + futures::pin_mut!(forward); + assert!( + futures::poll!(&mut forward).is_pending(), + "the WebSocket exited before the outbound router drained" + ); + + emit.notify_one(); + assert!( + !timeout(Duration::from_secs(1), forward) + .await + .expect("WebSocket did not close after the outbound router drained"), + "the closed agent channel should end the WebSocket loop" + ); + + let WsMessage::Text(text) = ws_rx.next().await.unwrap() else { + panic!("expected queued final text frame"); + }; + let message = serde_json::from_str::(&text).unwrap(); + assert!(matches!( + message, + RawJsonRpcMessage::Notification(notification) + if notification.method.as_ref() == "test/final" + )); + + registry.remove(&connection_id).await; + connection.shutdown().await; + } + #[tokio::test] async fn malformed_ws_frame_returns_parse_error_response_and_continues() { let (forwarded_tx, mut forwarded_rx) = mpsc::unbounded_channel(); From 946272daefa66c5f738abef8034a320acce90de1 Mon Sep 17 00:00:00 2001 From: Ben Brandt Date: Thu, 23 Jul 2026 15:08:21 +0200 Subject: [PATCH 3/3] docs(http): document transport lifecycle fixes --- src/agent-client-protocol-http/CHANGELOG.md | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/src/agent-client-protocol-http/CHANGELOG.md b/src/agent-client-protocol-http/CHANGELOG.md index af447dc..d0ec516 100644 --- a/src/agent-client-protocol-http/CHANGELOG.md +++ b/src/agent-client-protocol-http/CHANGELOG.md @@ -7,7 +7,11 @@ - **Breaking:** Upgrade to `agent-client-protocol` 2.x. Transport implementations and the core handlers/types they connect must be migrated together. - **Fixed:** Preserve incoming JSON-RPC batch frames and grouped responses across HTTP and WebSocket - transports, including session-aware HTTP routing. + transports, including session-aware HTTP routing. HTTP clients now validate and track every batch + entry and open streams returned by grouped `session/new` and `session/fork` responses. +- **Fixed:** Establish connection and session SSE streams before posting dependent messages while + continuing to deliver callbacks and complete earlier POSTs during setup. Pending setup is + cancelled cleanly when the caller exits. - **Fixed:** Keep call-bearing and invalid-request batch POSTs in peer order while allowing response-only frames, including malformed response-shaped values, to bypass a pending request when completing an SSE callback. @@ -15,6 +19,8 @@ to create a connection and return its success or rejection as one grouped JSON-RPC response, ignoring any leading response-only entries and buffering sibling side traffic for the connection stream until initialization completes. +- **Fixed:** Drain final routed messages to established HTTP SSE and WebSocket streams before + closing them when the connected agent exits. ## [1.3.0](https://github.com/agentclientprotocol/rust-sdk/compare/agent-client-protocol-http-v1.2.0...agent-client-protocol-http-v1.3.0) - 2026-07-20