Skip to main content

agent_client_protocol_http/
client.rs

1use std::{
2    collections::{HashMap, HashSet, VecDeque},
3    sync::{Arc, Mutex as StdMutex},
4};
5
6use agent_client_protocol::{
7    Agent, Channel, Client, ConnectTo, Error as AcpError, RawJsonRpcMessage, TransportBatchEntry,
8    TransportFrame,
9    schema::v1::{RequestId, Response as RpcResponse},
10};
11use async_tungstenite::tungstenite::Message as WsMessage;
12use futures::{
13    Stream, StreamExt,
14    channel::mpsc::{self, UnboundedSender},
15    future::{BoxFuture, FutureExt},
16    pin_mut,
17    stream::FuturesUnordered,
18};
19use thiserror::Error;
20use tracing::{debug, error, trace, warn};
21
22use crate::protocol::{
23    HEADER_CONNECTION_ID, HEADER_SESSION_ID, is_initialize_request, is_response_only_shape,
24    method_for_message, method_requires_session_header, session_id_from_message,
25};
26
27#[derive(Debug, Error)]
28pub enum HttpClientError {
29    #[error("invalid URL: {0}")]
30    InvalidUrl(#[from] url::ParseError),
31    #[error("failed to build HTTP client: {0}")]
32    Reqwest(#[from] reqwest::Error),
33}
34
35pub struct HttpClient {
36    endpoint: url::Url,
37    http: reqwest::Client,
38}
39
40impl std::fmt::Debug for HttpClient {
41    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
42        f.debug_struct("HttpClient")
43            .field("endpoint", &self.endpoint.as_str())
44            .finish_non_exhaustive()
45    }
46}
47
48impl HttpClient {
49    /// Create a client from a base URL and target the standard ACP endpoint.
50    ///
51    /// If the URL path is empty, `/acp` is used. Otherwise `/acp` is appended
52    /// unless the path already ends with `/acp`.
53    pub fn new(base_url: impl AsRef<str>) -> Result<Self, HttpClientError> {
54        Self::with_client(base_url, reqwest::Client::new())
55    }
56
57    /// Create a client that targets the exact endpoint URL.
58    ///
59    /// Use this when connecting to a server configured with a custom
60    /// `ServerOptions::path`.
61    pub fn with_endpoint(endpoint: impl AsRef<str>) -> Result<Self, HttpClientError> {
62        Self::with_endpoint_and_client(endpoint, reqwest::Client::new())
63    }
64
65    /// Create a client with a custom HTTP client and the standard ACP endpoint.
66    ///
67    /// If the URL path is empty, `/acp` is used. Otherwise `/acp` is appended
68    /// unless the path already ends with `/acp`.
69    pub fn with_client(
70        base_url: impl AsRef<str>,
71        http: reqwest::Client,
72    ) -> Result<Self, HttpClientError> {
73        let mut endpoint = url::Url::parse(base_url.as_ref())?;
74        let path = endpoint.path().trim_end_matches('/').to_string();
75        let path = if path.is_empty() {
76            "/acp".to_string()
77        } else if path.ends_with("/acp") {
78            path
79        } else {
80            format!("{path}/acp")
81        };
82        endpoint.set_path(&path);
83        Ok(Self { endpoint, http })
84    }
85
86    /// Create a client with a custom HTTP client and exact endpoint URL.
87    ///
88    /// Use this when connecting to a server configured with a custom
89    /// `ServerOptions::path`.
90    pub fn with_endpoint_and_client(
91        endpoint: impl AsRef<str>,
92        http: reqwest::Client,
93    ) -> Result<Self, HttpClientError> {
94        let endpoint = url::Url::parse(endpoint.as_ref())?;
95        Ok(Self { endpoint, http })
96    }
97
98    fn is_websocket(&self) -> bool {
99        matches!(self.endpoint.scheme(), "ws" | "wss")
100    }
101}
102
103impl ConnectTo<Client> for HttpClient {
104    async fn connect_to(self, client: impl ConnectTo<Agent>) -> Result<(), AcpError> {
105        let (channel, transport) = ConnectTo::<Client>::into_channel_and_future(self);
106        let shutdown_tx = channel.tx.clone();
107        match futures::future::select(
108            std::pin::pin!(client.connect_to(channel)),
109            std::pin::pin!(transport),
110        )
111        .await
112        {
113            futures::future::Either::Left((result, transport)) => {
114                result?;
115
116                // Reject sends from escaped client handles while preserving
117                // messages already accepted into the channel, then let the
118                // physical transport finish those messages.
119                shutdown_tx.close_channel();
120                transport.await
121            }
122            futures::future::Either::Right((result, _)) => result,
123        }
124    }
125
126    fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), AcpError>>) {
127        let (caller, transport) = Channel::duplex();
128        (caller, Box::pin(run(self, transport)))
129    }
130}
131
132async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> {
133    if client.is_websocket() {
134        return run_ws(client, channel).await;
135    }
136    let HttpClient { endpoint, http } = client;
137    let Channel {
138        rx: mut outgoing,
139        tx: incoming,
140    } = channel;
141    let (sse_event_tx, mut sse_event_rx) = mpsc::unbounded::<SseMessage>();
142    let connection = HttpConnection::new(endpoint, http);
143    let mut state = ClientState {
144        connection: connection.clone(),
145        open_session_streams: HashSet::new(),
146        pending_requests: HashMap::new(),
147        incoming,
148    };
149    let mut lifecycle = HttpTransportLifecycle::new(connection);
150    let mut posts = PostQueues::default();
151    let mut buffered_outgoing = VecDeque::new();
152    let mut outgoing_closed = false;
153
154    let result = 'transport: loop {
155        if outgoing_closed && buffered_outgoing.is_empty() && posts.is_empty() {
156            break Ok(());
157        }
158
159        let event = {
160            let outgoing_next = async {
161                if let Some(frame) = buffered_outgoing.pop_front() {
162                    Some(frame)
163                } else if outgoing_closed {
164                    futures::future::pending().await
165                } else {
166                    outgoing.next().await
167                }
168            }
169            .fuse();
170            let sse_event_next = sse_event_rx.next().fuse();
171            let sse_failure_next = lifecycle.next_sse_failure().fuse();
172            let ordered_post_next = posts.ordered.next_completion().fuse();
173            let response_post_next = posts.responses.next_completion().fuse();
174            pin_mut!(
175                outgoing_next,
176                sse_event_next,
177                sse_failure_next,
178                ordered_post_next,
179                response_post_next
180            );
181
182            futures::select! {
183                msg = outgoing_next => HttpLoopEvent::Outgoing(msg),
184                event = sse_event_next => HttpLoopEvent::SseEvent(event),
185                failure = sse_failure_next => HttpLoopEvent::SseFailure(failure),
186                post = ordered_post_next => HttpLoopEvent::Post(post),
187                post = response_post_next => HttpLoopEvent::Post(post),
188            }
189        };
190
191        let frame = match event {
192            HttpLoopEvent::Outgoing(msg) => {
193                let Some(frame) = msg else {
194                    outgoing_closed = true;
195                    continue;
196                };
197                frame
198            }
199            HttpLoopEvent::SseEvent(event) => {
200                let Some(event) = event else {
201                    continue;
202                };
203                let open_session_ids = state.sessions_to_open_for_responses(&event.frame);
204                state.deliver_frame(event.frame);
205                for session_id in open_session_ids {
206                    match lifecycle
207                        .start_sse(
208                            Some(session_id),
209                            sse_event_tx.clone(),
210                            SseStartContext {
211                                events: &mut sse_event_rx,
212                                outgoing: &mut outgoing,
213                                buffered_outgoing: &mut buffered_outgoing,
214                                posts: &mut posts,
215                                state: &mut state,
216                            },
217                        )
218                        .await
219                    {
220                        Ok(SseStartOutcome::Established) => {}
221                        Ok(SseStartOutcome::OutgoingClosed)
222                            if buffered_outgoing.is_empty() && posts.is_empty() =>
223                        {
224                            break 'transport Ok(());
225                        }
226                        Ok(SseStartOutcome::OutgoingClosed) => {
227                            break 'transport Err(sse_setup_blocked_output_error());
228                        }
229                        Err(error) => break 'transport Err(error),
230                    }
231                }
232                continue;
233            }
234            HttpLoopEvent::SseFailure(failure) => {
235                break Err(sse_failure_error(failure));
236            }
237            HttpLoopEvent::Post(completed) => {
238                if let Err(error) = handle_completed_post(&mut state, completed) {
239                    break Err(error);
240                }
241                continue;
242            }
243        };
244
245        let is_response_only = is_response_only_frame(&frame);
246        let msg = match frame {
247            TransportFrame::Single(message) => message,
248            frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => {
249                if state.connection.connection_id().is_none() {
250                    break Err(AcpError::invalid_request()
251                        .data("ACP HTTP transport: first message must be `initialize`"));
252                }
253                match state.prepare_frame_post(frame) {
254                    // Response-only batches answer SSE-delivered callbacks and
255                    // must not be blocked behind the request they answer.
256                    Ok((post, session_ids)) => {
257                        for session_id in session_ids {
258                            match lifecycle
259                                .start_sse(
260                                    Some(session_id),
261                                    sse_event_tx.clone(),
262                                    SseStartContext {
263                                        events: &mut sse_event_rx,
264                                        outgoing: &mut outgoing,
265                                        buffered_outgoing: &mut buffered_outgoing,
266                                        posts: &mut posts,
267                                        state: &mut state,
268                                    },
269                                )
270                                .await
271                            {
272                                Ok(SseStartOutcome::Established) => {}
273                                Ok(SseStartOutcome::OutgoingClosed) => {
274                                    break 'transport Err(sse_setup_blocked_output_error());
275                                }
276                                Err(error) => break 'transport Err(error),
277                            }
278                        }
279                        if is_response_only {
280                            posts.responses.push(post);
281                        } else {
282                            posts.ordered.push(post);
283                        }
284                    }
285                    Err(error) => {
286                        error!("POST failed: {error}");
287                        break Err(AcpError::internal_error().data(format!("POST: {error}")));
288                    }
289                }
290                continue;
291            }
292        };
293
294        if state.connection.connection_id().is_none() {
295            if !is_initialize_request(&msg) {
296                break Err(AcpError::invalid_request()
297                    .data("ACP HTTP transport: first message must be `initialize`"));
298            }
299            match state.initialize(msg).await {
300                Ok(InitializeOutcome::Connected) => {
301                    match lifecycle
302                        .start_sse(
303                            None,
304                            sse_event_tx.clone(),
305                            SseStartContext {
306                                events: &mut sse_event_rx,
307                                outgoing: &mut outgoing,
308                                buffered_outgoing: &mut buffered_outgoing,
309                                posts: &mut posts,
310                                state: &mut state,
311                            },
312                        )
313                        .await
314                    {
315                        Ok(SseStartOutcome::Established) => {}
316                        Ok(SseStartOutcome::OutgoingClosed) if buffered_outgoing.is_empty() => {
317                            break 'transport Ok(());
318                        }
319                        Ok(SseStartOutcome::OutgoingClosed) => {
320                            break 'transport Err(sse_setup_blocked_output_error());
321                        }
322                        Err(error) => break 'transport Err(error),
323                    }
324                }
325                Ok(InitializeOutcome::Rejected) => {}
326                Err(e) => {
327                    error!("initialize failed: {e}");
328                    break Err(AcpError::internal_error().data(format!("initialize: {e}")));
329                }
330            }
331            continue;
332        }
333
334        if let Some(session_id) = session_id_from_message(&msg) {
335            for session_id in state.register_session_streams([session_id]) {
336                match lifecycle
337                    .start_sse(
338                        Some(session_id),
339                        sse_event_tx.clone(),
340                        SseStartContext {
341                            events: &mut sse_event_rx,
342                            outgoing: &mut outgoing,
343                            buffered_outgoing: &mut buffered_outgoing,
344                            posts: &mut posts,
345                            state: &mut state,
346                        },
347                    )
348                    .await
349                {
350                    Ok(SseStartOutcome::Established) => {}
351                    Ok(SseStartOutcome::OutgoingClosed) => {
352                        break 'transport Err(sse_setup_blocked_output_error());
353                    }
354                    Err(error) => break 'transport Err(error),
355                }
356            }
357        }
358
359        match state.prepare_post(msg) {
360            // Responses answer SSE-delivered callbacks and must not be blocked
361            // behind a POST that may be waiting for that callback response.
362            Ok(post) if is_response_only => posts.responses.push(post),
363            Ok(post) => posts.ordered.push(post),
364            Err(e) => {
365                error!("POST failed: {e}");
366                break Err(AcpError::internal_error().data(format!("POST: {e}")));
367            }
368        }
369    };
370
371    lifecycle.close().await;
372    result
373}
374
375fn sse_failure_error(failure: SseFailure) -> AcpError {
376    let scope = failure.session_id.as_deref().unwrap_or("connection");
377    error!(session_id = ?failure.session_id, error = %failure.error, "SSE stream ended");
378    AcpError::internal_error().data(format!("{scope} SSE stream ended: {}", failure.error))
379}
380
381fn sse_setup_blocked_output_error() -> AcpError {
382    AcpError::internal_error()
383        .data("outgoing channel closed while accepted messages awaited SSE stream establishment")
384}
385
386fn handle_completed_post(
387    state: &mut ClientState,
388    completed: CompletedPost,
389) -> Result<(), AcpError> {
390    let CompletedPost {
391        pending_requests,
392        result,
393    } = completed;
394    if let Err(error) = result {
395        state.remove_pending_requests(&pending_requests);
396        error!("POST failed: {error}");
397        Err(AcpError::internal_error().data(format!("POST: {error}")))
398    } else {
399        Ok(())
400    }
401}
402
403fn queue_response_post(
404    state: &mut ClientState,
405    posts: &mut PostQueues,
406    frame: TransportFrame,
407) -> Result<(), AcpError> {
408    let post = match frame {
409        TransportFrame::Single(message) => state.prepare_post(message),
410        frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => {
411            state.prepare_frame_post(frame).map(|(post, session_ids)| {
412                debug_assert!(session_ids.is_empty());
413                post
414            })
415        }
416    }
417    .map_err(|error| {
418        error!("POST failed: {error}");
419        AcpError::internal_error().data(format!("POST: {error}"))
420    })?;
421    posts.responses.push(post);
422    Ok(())
423}
424
425fn is_response_only_frame(frame: &TransportFrame) -> bool {
426    match frame {
427        TransportFrame::Single(RawJsonRpcMessage::Response(_)) => true,
428        TransportFrame::Batch(batch) => batch.entries().all(|entry| match entry {
429            TransportBatchEntry::Message(RawJsonRpcMessage::Response(_)) => true,
430            TransportBatchEntry::Malformed { raw, .. } => is_response_only_shape(raw),
431            TransportBatchEntry::Message(
432                RawJsonRpcMessage::Request(_) | RawJsonRpcMessage::Notification(_),
433            ) => false,
434        }),
435        TransportFrame::Malformed { raw, .. } => {
436            serde_json::from_str(raw).is_ok_and(|value| is_response_only_shape(&value))
437        }
438        TransportFrame::Single(
439            RawJsonRpcMessage::Request(_) | RawJsonRpcMessage::Notification(_),
440        ) => false,
441    }
442}
443
444enum HttpLoopEvent {
445    Outgoing(Option<TransportFrame>),
446    SseEvent(Option<SseMessage>),
447    SseFailure(SseFailure),
448    Post(CompletedPost),
449}
450
451#[derive(Debug)]
452struct SseFailure {
453    session_id: Option<String>,
454    error: String,
455}
456
457#[derive(Debug)]
458struct SseMessage {
459    frame: TransportFrame,
460}
461
462#[derive(Clone, Debug)]
463struct HttpConnection {
464    endpoint: url::Url,
465    http: reqwest::Client,
466    connection_id: Arc<StdMutex<Option<String>>>,
467}
468
469impl HttpConnection {
470    fn new(endpoint: url::Url, http: reqwest::Client) -> Self {
471        Self {
472            endpoint,
473            http,
474            connection_id: Arc::new(StdMutex::new(None)),
475        }
476    }
477
478    fn post(&self) -> reqwest::RequestBuilder {
479        self.http.post(self.endpoint.clone())
480    }
481
482    fn get(&self) -> reqwest::RequestBuilder {
483        self.http.get(self.endpoint.clone())
484    }
485
486    fn set_connection_id(&self, connection_id: String) {
487        *self.connection_id.lock().expect("mutex poisoned") = Some(connection_id);
488    }
489
490    fn connection_id(&self) -> Option<String> {
491        self.connection_id.lock().expect("mutex poisoned").clone()
492    }
493
494    fn take_connection_id(&self) -> Option<String> {
495        self.connection_id.lock().expect("mutex poisoned").take()
496    }
497
498    fn clear_connection_id(&self, expected: &str) {
499        let mut connection_id = self.connection_id.lock().expect("mutex poisoned");
500        if connection_id.as_deref() == Some(expected) {
501            *connection_id = None;
502        }
503    }
504
505    async fn close(&self) {
506        let Some(connection_id) = self.connection_id() else {
507            return;
508        };
509        Self::send_close(
510            self.http.clone(),
511            self.endpoint.clone(),
512            connection_id.clone(),
513        )
514        .await;
515        self.clear_connection_id(&connection_id);
516    }
517
518    fn spawn_close(&self) {
519        let Some(connection_id) = self.take_connection_id() else {
520            return;
521        };
522        let http = self.http.clone();
523        let endpoint = self.endpoint.clone();
524        match tokio::runtime::Handle::try_current() {
525            Ok(handle) => {
526                drop(handle.spawn(Self::send_close(http, endpoint, connection_id)));
527            }
528            Err(e) => {
529                debug!("failed to spawn HTTP DELETE: {e}");
530            }
531        }
532    }
533
534    async fn send_close(http: reqwest::Client, endpoint: url::Url, connection_id: String) {
535        if let Err(e) = http
536            .delete(endpoint)
537            .header(HEADER_CONNECTION_ID, connection_id)
538            .send()
539            .await
540        {
541            debug!("DELETE failed (ignored): {e}");
542        }
543    }
544}
545
546#[derive(Debug)]
547struct HttpTransportLifecycle {
548    connection: HttpConnection,
549    sse_tasks: SseTasks,
550}
551
552#[derive(Clone, Copy, Debug, Eq, PartialEq)]
553enum SseStartOutcome {
554    Established,
555    OutgoingClosed,
556}
557
558struct SseStartContext<'a> {
559    events: &'a mut mpsc::UnboundedReceiver<SseMessage>,
560    outgoing: &'a mut mpsc::UnboundedReceiver<TransportFrame>,
561    buffered_outgoing: &'a mut VecDeque<TransportFrame>,
562    posts: &'a mut PostQueues,
563    state: &'a mut ClientState,
564}
565
566impl HttpTransportLifecycle {
567    fn new(connection: HttpConnection) -> Self {
568        Self {
569            connection,
570            sse_tasks: SseTasks::default(),
571        }
572    }
573
574    async fn start_sse(
575        &mut self,
576        session_id: Option<String>,
577        event_tx: UnboundedSender<SseMessage>,
578        context: SseStartContext<'_>,
579    ) -> Result<SseStartOutcome, AcpError> {
580        let SseStartContext {
581            events,
582            outgoing,
583            buffered_outgoing,
584            posts,
585            state,
586        } = context;
587        let mut establishing = FuturesUnordered::new();
588        establishing.push(self.begin_sse(session_id, event_tx.clone()));
589
590        loop {
591            if establishing.is_empty() {
592                return Ok(SseStartOutcome::Established);
593            }
594            let outcome = {
595                let failure = self.sse_tasks.next_failure().fuse();
596                let established_next = establishing.next().fuse();
597                let sse_event_next = events.next().fuse();
598                let outgoing_next = outgoing.next().fuse();
599                let ordered_post_next = posts.ordered.next_completion().fuse();
600                let response_post_next = posts.responses.next_completion().fuse();
601                pin_mut!(
602                    failure,
603                    established_next,
604                    sse_event_next,
605                    outgoing_next,
606                    ordered_post_next,
607                    response_post_next
608                );
609                futures::select_biased! {
610                    failure = failure => SseStartWait::Failure(failure),
611                    established = established_next => SseStartWait::Established(established),
612                    event = sse_event_next => SseStartWait::SseEvent(event),
613                    post = response_post_next => SseStartWait::Post(post),
614                    post = ordered_post_next => SseStartWait::Post(post),
615                    outgoing = outgoing_next => SseStartWait::Outgoing(outgoing),
616                }
617            };
618            match outcome {
619                SseStartWait::Established(Some(Ok(()))) => {}
620                SseStartWait::Established(Some(Err(_))) => {
621                    return Err(sse_failure_error(self.sse_tasks.next_failure().await));
622                }
623                SseStartWait::Established(None) => {
624                    return Ok(SseStartOutcome::Established);
625                }
626                SseStartWait::Failure(failure) => return Err(sse_failure_error(failure)),
627                SseStartWait::SseEvent(Some(event)) => {
628                    let open_session_ids = state.sessions_to_open_for_responses(&event.frame);
629                    state.deliver_frame(event.frame);
630                    for session_id in open_session_ids {
631                        establishing.push(self.begin_sse(Some(session_id), event_tx.clone()));
632                    }
633                }
634                SseStartWait::SseEvent(None) => {
635                    return Err(AcpError::internal_error().data("SSE event channel closed"));
636                }
637                SseStartWait::Post(completed) => handle_completed_post(state, completed)?,
638                SseStartWait::Outgoing(Some(frame)) if is_response_only_frame(&frame) => {
639                    queue_response_post(state, posts, frame)?;
640                }
641                SseStartWait::Outgoing(Some(frame)) => buffered_outgoing.push_back(frame),
642                SseStartWait::Outgoing(None) => return Ok(SseStartOutcome::OutgoingClosed),
643            }
644        }
645    }
646
647    fn begin_sse(
648        &mut self,
649        session_id: Option<String>,
650        event_tx: UnboundedSender<SseMessage>,
651    ) -> futures::channel::oneshot::Receiver<()> {
652        let (established_tx, established_rx) = futures::channel::oneshot::channel();
653        self.sse_tasks.push(run_sse(
654            self.connection.clone(),
655            session_id,
656            event_tx,
657            established_tx,
658        ));
659        established_rx
660    }
661
662    async fn next_sse_failure(&mut self) -> SseFailure {
663        self.sse_tasks.next_failure().await
664    }
665
666    async fn close(&mut self) {
667        self.connection.close().await;
668        self.sse_tasks.abort_all();
669    }
670}
671
672enum SseStartWait {
673    Established(Option<Result<(), futures::channel::oneshot::Canceled>>),
674    Failure(SseFailure),
675    SseEvent(Option<SseMessage>),
676    Post(CompletedPost),
677    Outgoing(Option<TransportFrame>),
678}
679
680impl Drop for HttpTransportLifecycle {
681    fn drop(&mut self) {
682        self.sse_tasks.abort_all();
683        self.connection.spawn_close();
684    }
685}
686
687fn run_sse(
688    connection: HttpConnection,
689    session_id: Option<String>,
690    event_tx: UnboundedSender<SseMessage>,
691    established_tx: futures::channel::oneshot::Sender<()>,
692) -> BoxFuture<'static, SseFailure> {
693    Box::pin(async move {
694        let label = session_id.clone();
695        let error = match read_sse(connection, session_id, event_tx, established_tx).await {
696            Ok(()) => "SSE stream closed".to_string(),
697            Err(e) => e,
698        };
699        warn!(session_id = ?label, "SSE stream ended: {error}");
700        SseFailure {
701            session_id: label,
702            error,
703        }
704    })
705}
706
707#[derive(Debug, Default)]
708struct SseTasks {
709    handles: FuturesUnordered<BoxFuture<'static, SseFailure>>,
710}
711
712impl SseTasks {
713    fn push(&mut self, task: BoxFuture<'static, SseFailure>) {
714        self.handles.push(task);
715    }
716
717    async fn next_failure(&mut self) -> SseFailure {
718        loop {
719            if let Some(failure) = self.handles.next().await {
720                return failure;
721            }
722            futures::future::pending::<()>().await;
723        }
724    }
725
726    fn abort_all(&mut self) {
727        self.handles = FuturesUnordered::new();
728    }
729}
730
731struct ClientState {
732    connection: HttpConnection,
733    open_session_streams: HashSet<String>,
734    pending_requests: HashMap<RequestId, VecDeque<String>>,
735    incoming: futures::channel::mpsc::UnboundedSender<TransportFrame>,
736}
737
738struct PendingPost {
739    pending_requests: Vec<(RequestId, String)>,
740    response: BoxFuture<'static, Result<(), String>>,
741}
742
743impl PendingPost {
744    fn into_completion(self) -> BoxFuture<'static, CompletedPost> {
745        let Self {
746            pending_requests,
747            response,
748        } = self;
749        async move {
750            CompletedPost {
751                pending_requests,
752                result: response.await,
753            }
754        }
755        .boxed()
756    }
757}
758
759#[derive(Debug)]
760struct CompletedPost {
761    pending_requests: Vec<(RequestId, String)>,
762    result: Result<(), String>,
763}
764
765#[derive(Default)]
766struct PostQueue {
767    queued: VecDeque<PendingPost>,
768    in_flight: Option<BoxFuture<'static, CompletedPost>>,
769}
770
771#[derive(Default)]
772struct PostQueues {
773    ordered: PostQueue,
774    responses: PostQueue,
775}
776
777impl PostQueues {
778    fn is_empty(&self) -> bool {
779        self.ordered.is_empty() && self.responses.is_empty()
780    }
781}
782
783impl PostQueue {
784    fn push(&mut self, post: PendingPost) {
785        self.queued.push_back(post);
786        self.start_next();
787    }
788
789    async fn next_completion(&mut self) -> CompletedPost {
790        loop {
791            self.start_next();
792            if let Some(in_flight) = self.in_flight.as_mut() {
793                let completed = in_flight.await;
794                self.in_flight = None;
795                return completed;
796            }
797            futures::future::pending::<()>().await;
798        }
799    }
800
801    fn start_next(&mut self) {
802        if self.in_flight.is_none()
803            && let Some(post) = self.queued.pop_front()
804        {
805            self.in_flight = Some(post.into_completion());
806        }
807    }
808
809    fn is_empty(&self) -> bool {
810        self.queued.is_empty() && self.in_flight.is_none()
811    }
812}
813
814#[derive(Clone, Copy, Debug, Eq, PartialEq)]
815enum InitializeOutcome {
816    Connected,
817    Rejected,
818}
819
820impl ClientState {
821    async fn initialize(&self, msg: RawJsonRpcMessage) -> Result<InitializeOutcome, String> {
822        let response = self
823            .connection
824            .post()
825            .header("Content-Type", "application/json")
826            .header("Accept", "application/json")
827            .json(&msg)
828            .send()
829            .await
830            .map_err(|e| e.to_string())?;
831
832        let connection_id = response
833            .headers()
834            .get(HEADER_CONNECTION_ID)
835            .and_then(|v| v.to_str().ok())
836            .map(String::from);
837        if let Some(connection_id) = &connection_id {
838            self.connection.set_connection_id(connection_id.clone());
839        }
840
841        if !response.status().is_success() {
842            let status = response.status();
843            let body = response.text().await.unwrap_or_default();
844            return Err(format!("HTTP {status}: {body}"));
845        }
846
847        let body = response.text().await.map_err(|error| error.to_string())?;
848        let message = match TransportFrame::parse_json(&body) {
849            TransportFrame::Single(message) => message,
850            TransportFrame::Malformed { error, .. } => {
851                return Err(format!("invalid initialize response: {error}"));
852            }
853            TransportFrame::Batch(_) => {
854                return Err("initialize response must not be a JSON-RPC batch".to_string());
855            }
856        };
857
858        if matches!(
859            message,
860            RawJsonRpcMessage::Response(RpcResponse::Error { .. })
861        ) {
862            self.deliver(message);
863            self.connection.close().await;
864            return Ok(InitializeOutcome::Rejected);
865        }
866
867        connection_id
868            .ok_or_else(|| format!("server did not return {HEADER_CONNECTION_ID} header"))?;
869        self.deliver(message);
870        Ok(InitializeOutcome::Connected)
871    }
872
873    fn prepare_post(&mut self, msg: RawJsonRpcMessage) -> Result<PendingPost, String> {
874        let session_id = validated_session_id(&msg)?;
875        let connection_id = self
876            .connection
877            .connection_id()
878            .ok_or_else(|| "POST attempted before initialize".to_string())?;
879        let mut request = self
880            .connection
881            .post()
882            .header("Accept", "application/json")
883            .header(HEADER_CONNECTION_ID, connection_id)
884            .json(&msg);
885        if let Some(session_id) = session_id {
886            request = request.header(HEADER_SESSION_ID, session_id);
887        }
888
889        let pending_requests = pending_request_for_message(&msg)
890            .into_iter()
891            .collect::<Vec<_>>();
892        self.track_pending_requests(&pending_requests);
893
894        let response = async move {
895            let response = request.send().await.map_err(|e| e.to_string())?;
896            if response.status().as_u16() != 202 && !response.status().is_success() {
897                let status = response.status();
898                let body = response.text().await.unwrap_or_default();
899                return Err(format!("HTTP {status}: {body}"));
900            }
901            Ok(())
902        };
903        Ok(PendingPost {
904            pending_requests,
905            response: response.boxed(),
906        })
907    }
908
909    fn prepare_frame_post(
910        &mut self,
911        frame: TransportFrame,
912    ) -> Result<(PendingPost, Vec<String>), String> {
913        let bookkeeping = FrameBookkeeping::for_frame(&frame)?;
914        let connection_id = self
915            .connection
916            .connection_id()
917            .ok_or_else(|| "POST attempted before initialize".to_string())?;
918        let body = frame.to_json().map_err(|error| error.to_string())?;
919        let request = self
920            .connection
921            .post()
922            .header("Content-Type", "application/json")
923            .header("Accept", "application/json")
924            .header(HEADER_CONNECTION_ID, connection_id)
925            .body(body);
926        let response = async move {
927            let response = request.send().await.map_err(|error| error.to_string())?;
928            if response.status().as_u16() != 202 && !response.status().is_success() {
929                let status = response.status();
930                let body = response.text().await.unwrap_or_default();
931                return Err(format!("HTTP {status}: {body}"));
932            }
933            Ok(())
934        };
935        self.track_pending_requests(&bookkeeping.pending_requests);
936        let session_ids = self.register_session_streams(bookkeeping.session_ids);
937        Ok((
938            PendingPost {
939                pending_requests: bookkeeping.pending_requests,
940                response: response.boxed(),
941            },
942            session_ids,
943        ))
944    }
945
946    fn track_pending_requests(&mut self, pending_requests: &[(RequestId, String)]) {
947        for (id, method) in pending_requests {
948            self.pending_requests
949                .entry(id.clone())
950                .or_default()
951                .push_back(method.clone());
952        }
953    }
954
955    fn remove_pending_requests(&mut self, pending_requests: &[(RequestId, String)]) {
956        for (id, method) in pending_requests.iter().rev() {
957            let remove_entry = self.pending_requests.get_mut(id).is_some_and(|methods| {
958                if let Some(index) = methods.iter().rposition(|candidate| candidate == method) {
959                    methods.remove(index);
960                }
961                methods.is_empty()
962            });
963            if remove_entry {
964                self.pending_requests.remove(id);
965            }
966        }
967    }
968
969    fn take_pending_request_method(&mut self, id: &RequestId) -> Option<String> {
970        let (method, remove_entry) = {
971            let methods = self.pending_requests.get_mut(id)?;
972            (methods.pop_front(), methods.is_empty())
973        };
974        if remove_entry {
975            self.pending_requests.remove(id);
976        }
977        method
978    }
979
980    fn register_session_streams(
981        &mut self,
982        session_ids: impl IntoIterator<Item = String>,
983    ) -> Vec<String> {
984        session_ids
985            .into_iter()
986            .filter(|session_id| self.open_session_streams.insert(session_id.clone()))
987            .collect()
988    }
989
990    fn sessions_to_open_for_responses(&mut self, frame: &TransportFrame) -> Vec<String> {
991        match frame {
992            TransportFrame::Single(message) => self
993                .session_to_open_for_response(message)
994                .into_iter()
995                .collect(),
996            TransportFrame::Batch(batch) => batch
997                .entries()
998                .filter_map(|entry| match entry {
999                    TransportBatchEntry::Message(message) => {
1000                        self.session_to_open_for_response(message)
1001                    }
1002                    TransportBatchEntry::Malformed { .. } => None,
1003                })
1004                .collect(),
1005            TransportFrame::Malformed { .. } => Vec::new(),
1006        }
1007    }
1008
1009    fn session_to_open_for_response(&mut self, msg: &RawJsonRpcMessage) -> Option<String> {
1010        let RawJsonRpcMessage::Response(response) = msg else {
1011            return None;
1012        };
1013        let id = msg.response_id().and_then(pending_request_key)?;
1014        let method = self.take_pending_request_method(&id);
1015
1016        if !method.as_deref().is_some_and(is_session_opening_method) {
1017            return None;
1018        }
1019        let RpcResponse::Result { result, .. } = response else {
1020            return None;
1021        };
1022        let session_id = result
1023            .get("sessionId")
1024            .and_then(|v| v.as_str())
1025            .map(String::from)?;
1026
1027        if self.open_session_streams.insert(session_id.clone()) {
1028            Some(session_id)
1029        } else {
1030            None
1031        }
1032    }
1033
1034    fn deliver(&self, msg: RawJsonRpcMessage) {
1035        self.deliver_frame(TransportFrame::Single(msg));
1036    }
1037
1038    fn deliver_frame(&self, frame: TransportFrame) {
1039        if self.incoming.unbounded_send(frame).is_err() {
1040            debug!("upstream channel closed; dropping inbound message");
1041        }
1042    }
1043}
1044
1045#[derive(Default)]
1046struct FrameBookkeeping {
1047    session_ids: Vec<String>,
1048    pending_requests: Vec<(RequestId, String)>,
1049}
1050
1051impl FrameBookkeeping {
1052    fn for_frame(frame: &TransportFrame) -> Result<Self, String> {
1053        let mut bookkeeping = Self::default();
1054        match frame {
1055            TransportFrame::Single(message) => bookkeeping.add_message(message)?,
1056            TransportFrame::Batch(batch) => {
1057                for entry in batch.entries() {
1058                    if let TransportBatchEntry::Message(message) = entry {
1059                        bookkeeping.add_message(message)?;
1060                    }
1061                }
1062            }
1063            TransportFrame::Malformed { .. } => {}
1064        }
1065        Ok(bookkeeping)
1066    }
1067
1068    fn add_message(&mut self, message: &RawJsonRpcMessage) -> Result<(), String> {
1069        if let Some(session_id) = validated_session_id(message)?
1070            && !self.session_ids.contains(&session_id)
1071        {
1072            self.session_ids.push(session_id);
1073        }
1074        if let Some(pending_request) = pending_request_for_message(message) {
1075            self.pending_requests.push(pending_request);
1076        }
1077        Ok(())
1078    }
1079}
1080
1081fn validated_session_id(msg: &RawJsonRpcMessage) -> Result<Option<String>, String> {
1082    let Some(method) = method_for_message(msg) else {
1083        return Ok(None);
1084    };
1085    let session_id = session_id_from_message(msg);
1086    if method_requires_session_header(method) && session_id.is_none() {
1087        return Err(format!("method `{method}` requires sessionId in params"));
1088    }
1089    Ok(session_id)
1090}
1091
1092fn is_session_opening_method(method: &str) -> bool {
1093    matches!(method, "session/new" | "session/fork")
1094}
1095
1096async fn read_sse(
1097    connection: HttpConnection,
1098    session_id: Option<String>,
1099    event_tx: UnboundedSender<SseMessage>,
1100    established_tx: futures::channel::oneshot::Sender<()>,
1101) -> Result<(), String> {
1102    let connection_id = connection
1103        .connection_id()
1104        .ok_or_else(|| "SSE attempted before initialize".to_string())?;
1105    let mut request = connection
1106        .get()
1107        .header("Accept", "text/event-stream")
1108        .header(HEADER_CONNECTION_ID, connection_id);
1109    if let Some(session_id) = &session_id {
1110        request = request.header(HEADER_SESSION_ID, session_id);
1111    }
1112
1113    let response = request.send().await.map_err(|e| e.to_string())?;
1114    if !response.status().is_success() {
1115        return Err(format!("HTTP {}", response.status()));
1116    }
1117    trace!(session_id = ?session_id, "SSE stream open");
1118    let _ = established_tx.send(());
1119
1120    let mut events = eventsource_stream::EventStream::new(response.bytes_stream());
1121    while let Some(event) = events.next().await {
1122        let event = event.map_err(|e| e.to_string())?;
1123        let payload = event.data;
1124        if payload.is_empty() {
1125            continue;
1126        }
1127        let frame = TransportFrame::parse_json(&payload);
1128
1129        if event_tx.unbounded_send(SseMessage { frame }).is_err() {
1130            return Err("upstream channel closed".to_string());
1131        }
1132    }
1133    Ok(())
1134}
1135
1136fn pending_request_for_message(msg: &RawJsonRpcMessage) -> Option<(RequestId, String)> {
1137    let RawJsonRpcMessage::Request(request) = msg else {
1138        return None;
1139    };
1140    pending_request_key(&request.id).map(|id| (id, request.method.to_string()))
1141}
1142
1143fn pending_request_key(id: &RequestId) -> Option<RequestId> {
1144    match id {
1145        RequestId::Null => None,
1146        RequestId::Number(_) | RequestId::Str(_) => Some(id.clone()),
1147    }
1148}
1149
1150async fn run_ws(client: HttpClient, channel: Channel) -> Result<(), AcpError> {
1151    let HttpClient { endpoint, .. } = client;
1152
1153    let (ws_stream, response) = async_tungstenite::tokio::connect_async(endpoint.as_str())
1154        .await
1155        .map_err(|e| AcpError::internal_error().data(format!("WebSocket connect failed: {e}")))?;
1156    trace!(
1157        status = %response.status(),
1158        "WebSocket connection established"
1159    );
1160    let (ws_tx, ws_rx) = ws_stream.split();
1161
1162    drive_ws(ws_tx, ws_rx, channel).await
1163}
1164
1165trait WsSink {
1166    fn send(
1167        &mut self,
1168        message: WsMessage,
1169    ) -> impl std::future::Future<Output = Result<(), String>> + Send;
1170}
1171
1172impl<S> WsSink for async_tungstenite::WebSocketSender<S>
1173where
1174    S: futures::AsyncRead + futures::AsyncWrite + Unpin + Send,
1175{
1176    async fn send(&mut self, message: WsMessage) -> Result<(), String> {
1177        async_tungstenite::WebSocketSender::send(self, message)
1178            .await
1179            .map_err(|error| error.to_string())
1180    }
1181}
1182
1183async fn drive_ws<Tx, Rx, RxError>(
1184    mut ws_tx: Tx,
1185    mut ws_rx: Rx,
1186    channel: Channel,
1187) -> Result<(), AcpError>
1188where
1189    Tx: WsSink,
1190    Rx: Stream<Item = Result<WsMessage, RxError>> + Unpin,
1191    RxError: std::fmt::Display,
1192{
1193    let Channel {
1194        rx: mut outgoing,
1195        tx: incoming,
1196    } = channel;
1197    let writer = async move {
1198        while let Some(frame) = outgoing.next().await {
1199            let text = match frame.to_json() {
1200                Ok(text) => text,
1201                Err(error) => {
1202                    error!("failed to serialize outbound frame: {error}");
1203                    return Err(AcpError::internal_error().data(format!("serialize: {error}")));
1204                }
1205            };
1206            if let Err(error) = ws_tx.send(WsMessage::Text(text.into())).await {
1207                error!("WebSocket send failed: {error}");
1208                return Err(AcpError::internal_error().data(format!("ws send: {error}")));
1209            }
1210        }
1211
1212        drop(ws_tx.send(WsMessage::Close(None)).await);
1213        Ok(())
1214    };
1215
1216    let reader = async move {
1217        let mut discard_incoming = false;
1218        loop {
1219            match ws_rx.next().await {
1220                Some(Ok(WsMessage::Text(text))) => {
1221                    if discard_incoming {
1222                        continue;
1223                    }
1224                    let frame = TransportFrame::parse_json(text.as_str());
1225                    if incoming.unbounded_send(frame).is_err() {
1226                        debug!(
1227                            "upstream channel closed; discarding WS input while draining output"
1228                        );
1229                        discard_incoming = true;
1230                    }
1231                }
1232                Some(Ok(WsMessage::Binary(_))) => {
1233                    warn!("ignoring binary WebSocket frame (ACP uses text)");
1234                }
1235                Some(Ok(WsMessage::Ping(_) | WsMessage::Pong(_) | WsMessage::Frame(_))) => {}
1236                Some(Ok(WsMessage::Close(frame))) => {
1237                    debug!("server closed WebSocket: {frame:?}");
1238                    return Err(AcpError::internal_error()
1239                        .data(format!("WebSocket closed by peer: {frame:?}")));
1240                }
1241                Some(Err(e)) => {
1242                    error!("WebSocket receive error: {e}");
1243                    return Err(AcpError::internal_error().data(format!("ws recv: {e}")));
1244                }
1245                None => {
1246                    return Err(AcpError::internal_error().data("WebSocket stream ended"));
1247                }
1248            }
1249        }
1250    };
1251
1252    pin_mut!(writer, reader);
1253    match futures::future::select(writer, reader).await {
1254        futures::future::Either::Left((result, _))
1255        | futures::future::Either::Right((result, _)) => result,
1256    }
1257}
1258
1259#[cfg(test)]
1260mod tests {
1261    use std::{
1262        convert::Infallible,
1263        sync::{
1264            Arc,
1265            atomic::{AtomicBool, AtomicUsize, Ordering},
1266        },
1267        time::Duration,
1268    };
1269
1270    use agent_client_protocol::{TransportBatch, schema::v1::RequestId};
1271    use axum::{
1272        Json, Router,
1273        extract::{WebSocketUpgrade, ws::Message as AxumWsMessage},
1274        http::{HeaderMap, HeaderValue, StatusCode},
1275        response::{IntoResponse, Sse, sse::Event},
1276        routing::{get, post},
1277    };
1278    use serde_json::json;
1279    use tokio::{
1280        net::TcpListener,
1281        sync::Notify,
1282        time::{sleep, timeout},
1283    };
1284
1285    use super::*;
1286
1287    struct PostsThenExitClient {
1288        finish: Arc<Notify>,
1289        finished: Arc<Notify>,
1290        escaped_tx: futures::channel::oneshot::Sender<
1291            futures::channel::mpsc::UnboundedSender<TransportFrame>,
1292        >,
1293    }
1294
1295    struct InitializeThenExitClient {
1296        sse_started: Arc<Notify>,
1297        finished: Arc<Notify>,
1298    }
1299
1300    struct QueueOutgoingThenText {
1301        text: Option<WsMessage>,
1302        outgoing: Option<mpsc::UnboundedSender<TransportFrame>>,
1303    }
1304
1305    struct RecordingWsSink(mpsc::UnboundedSender<WsMessage>);
1306
1307    struct BackpressuredWsSink {
1308        output: mpsc::UnboundedSender<WsMessage>,
1309        started: mpsc::UnboundedSender<()>,
1310        release: Option<futures::channel::oneshot::Receiver<()>>,
1311    }
1312
1313    struct ReleaseBackpressureOnPoll {
1314        started: mpsc::UnboundedReceiver<()>,
1315        release: Option<futures::channel::oneshot::Sender<()>>,
1316    }
1317
1318    fn single_frame(message: RawJsonRpcMessage) -> TransportFrame {
1319        TransportFrame::Single(message)
1320    }
1321
1322    fn into_single_message(frame: TransportFrame) -> Result<RawJsonRpcMessage, AcpError> {
1323        match frame {
1324            TransportFrame::Single(message) => Ok(message),
1325            TransportFrame::Malformed { error, .. } => Err(error),
1326            TransportFrame::Batch(_) => {
1327                Err(AcpError::internal_error().data("expected one JSON-RPC message"))
1328            }
1329        }
1330    }
1331
1332    trait TransportFrameTestExt {
1333        fn unwrap(self) -> RawJsonRpcMessage;
1334    }
1335
1336    impl TransportFrameTestExt for TransportFrame {
1337        fn unwrap(self) -> RawJsonRpcMessage {
1338            into_single_message(self).unwrap()
1339        }
1340    }
1341
1342    #[test]
1343    fn malformed_response_shapes_bypass_only_when_the_whole_frame_is_response_only() {
1344        let standalone_response = TransportFrame::parse_json(
1345            r#"{"jsonrpc":"2.0","id":1,"result":{},"error":{"code":-32603}}"#,
1346        );
1347        assert!(is_response_only_frame(&standalone_response));
1348
1349        let response_batch = TransportFrame::parse_json(
1350            r#"[
1351                {"jsonrpc":"2.0","id":1,"result":{}},
1352                {"jsonrpc":"2.0","id":2,"result":{},"error":{"code":-32603}}
1353            ]"#,
1354        );
1355        assert!(is_response_only_frame(&response_batch));
1356
1357        let scalar_batch = TransportFrame::parse_json(
1358            r#"[
1359                {"jsonrpc":"2.0","id":1,"result":{}},
1360                17
1361            ]"#,
1362        );
1363        assert!(!is_response_only_frame(&scalar_batch));
1364
1365        let call_shaped = TransportFrame::parse_json(
1366            r#"{"jsonrpc":"2.0","id":1,"method":"custom/call","result":{}}"#,
1367        );
1368        assert!(!is_response_only_frame(&call_shaped));
1369    }
1370
1371    fn initialized_client_state() -> ClientState {
1372        let connection = HttpConnection::new(
1373            url::Url::parse("http://127.0.0.1/acp").unwrap(),
1374            reqwest::Client::new(),
1375        );
1376        connection.set_connection_id("connection-1".to_string());
1377        let (incoming, _incoming_rx) = mpsc::unbounded();
1378        ClientState {
1379            connection,
1380            open_session_streams: HashSet::new(),
1381            pending_requests: HashMap::new(),
1382            incoming,
1383        }
1384    }
1385
1386    #[test]
1387    fn batch_post_validation_happens_before_tracking_requests_or_sessions() {
1388        let mut state = initialized_client_state();
1389        let frame = TransportFrame::Batch(
1390            TransportBatch::from_messages([
1391                RawJsonRpcMessage::request(
1392                    "custom/valid".to_string(),
1393                    json!({}),
1394                    RequestId::Number(1),
1395                )
1396                .unwrap(),
1397                RawJsonRpcMessage::request(
1398                    "session/prompt".to_string(),
1399                    json!({ "prompt": [] }),
1400                    RequestId::Number(2),
1401                )
1402                .unwrap(),
1403            ])
1404            .unwrap(),
1405        );
1406
1407        let Err(error) = state.prepare_frame_post(frame) else {
1408            panic!("batch should require sessionId for session/prompt");
1409        };
1410
1411        assert_eq!(
1412            error,
1413            "method `session/prompt` requires sessionId in params"
1414        );
1415        assert!(state.pending_requests.is_empty());
1416        assert!(state.open_session_streams.is_empty());
1417    }
1418
1419    #[test]
1420    fn batch_post_tracks_every_non_null_request_and_rolls_back_from_the_back() {
1421        let mut state = initialized_client_state();
1422        state.track_pending_requests(&[(RequestId::Number(7), "session/fork".to_string())]);
1423        let frame = TransportFrame::Batch(
1424            TransportBatch::from_messages([
1425                RawJsonRpcMessage::request(
1426                    "session/fork".to_string(),
1427                    json!({ "sessionId": "source-a" }),
1428                    RequestId::Number(7),
1429                )
1430                .unwrap(),
1431                RawJsonRpcMessage::request(
1432                    "custom/request".to_string(),
1433                    json!({ "sessionId": "source-b" }),
1434                    RequestId::Number(7),
1435                )
1436                .unwrap(),
1437                RawJsonRpcMessage::request(
1438                    "session/fork".to_string(),
1439                    json!({ "sessionId": "source-a" }),
1440                    RequestId::Null,
1441                )
1442                .unwrap(),
1443            ])
1444            .unwrap(),
1445        );
1446
1447        let (post, session_ids) = state.prepare_frame_post(frame).unwrap();
1448
1449        assert_eq!(session_ids, ["source-a", "source-b"]);
1450        assert_eq!(
1451            state.pending_requests.get(&RequestId::Number(7)).unwrap(),
1452            &VecDeque::from([
1453                "session/fork".to_string(),
1454                "session/fork".to_string(),
1455                "custom/request".to_string(),
1456            ])
1457        );
1458        assert_eq!(
1459            post.pending_requests,
1460            [
1461                (RequestId::Number(7), "session/fork".to_string()),
1462                (RequestId::Number(7), "custom/request".to_string()),
1463            ]
1464        );
1465        assert!(!state.pending_requests.contains_key(&RequestId::Null));
1466
1467        state.remove_pending_requests(&post.pending_requests);
1468
1469        assert_eq!(
1470            state.pending_requests.get(&RequestId::Number(7)).unwrap(),
1471            &VecDeque::from(["session/fork".to_string()])
1472        );
1473    }
1474
1475    impl WsSink for RecordingWsSink {
1476        fn send(
1477            &mut self,
1478            message: WsMessage,
1479        ) -> impl std::future::Future<Output = Result<(), String>> + Send {
1480            std::future::ready(
1481                self.0
1482                    .unbounded_send(message)
1483                    .map_err(|error| error.to_string()),
1484            )
1485        }
1486    }
1487
1488    impl WsSink for BackpressuredWsSink {
1489        async fn send(&mut self, message: WsMessage) -> Result<(), String> {
1490            self.output
1491                .unbounded_send(message)
1492                .map_err(|error| error.to_string())?;
1493            if let Some(release) = self.release.take() {
1494                self.started
1495                    .unbounded_send(())
1496                    .map_err(|error| error.to_string())?;
1497                release
1498                    .await
1499                    .map_err(|_| "mock WebSocket reader did not release send".to_string())?;
1500            }
1501            Ok(())
1502        }
1503    }
1504
1505    impl Stream for QueueOutgoingThenText {
1506        type Item = Result<WsMessage, std::io::Error>;
1507
1508        fn poll_next(
1509            mut self: std::pin::Pin<&mut Self>,
1510            _cx: &mut std::task::Context<'_>,
1511        ) -> std::task::Poll<Option<Self::Item>> {
1512            // Make input ready immediately after queueing output. If the output
1513            // branch was polled first it was still empty, so either poll order
1514            // deterministically selects this input frame first.
1515            if let Some(outgoing) = self.outgoing.take() {
1516                for method in ["custom/first", "custom/second"] {
1517                    outgoing
1518                        .unbounded_send(single_frame(
1519                            RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(),
1520                        ))
1521                        .unwrap();
1522                }
1523            }
1524            if let Some(text) = self.text.take() {
1525                return std::task::Poll::Ready(Some(Ok(text)));
1526            }
1527            std::task::Poll::Pending
1528        }
1529    }
1530
1531    impl Stream for ReleaseBackpressureOnPoll {
1532        type Item = Result<WsMessage, std::io::Error>;
1533
1534        fn poll_next(
1535            mut self: std::pin::Pin<&mut Self>,
1536            cx: &mut std::task::Context<'_>,
1537        ) -> std::task::Poll<Option<Self::Item>> {
1538            if let std::task::Poll::Ready(Some(())) =
1539                std::pin::Pin::new(&mut self.started).poll_next(cx)
1540                && let Some(release) = self.release.take()
1541            {
1542                let _result = release.send(());
1543            }
1544            std::task::Poll::Pending
1545        }
1546    }
1547
1548    impl ConnectTo<Agent> for PostsThenExitClient {
1549        async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), AcpError> {
1550            let Self {
1551                finish,
1552                finished,
1553                escaped_tx,
1554            } = self;
1555            let (mut channel, transport) = agent.into_channel_and_future();
1556            let client = async move {
1557                escaped_tx.send(channel.tx.clone()).map_err(|_| {
1558                    AcpError::internal_error().data("escaped sender observer dropped")
1559                })?;
1560                channel
1561                    .tx
1562                    .unbounded_send(single_frame(
1563                        RawJsonRpcMessage::request(
1564                            "initialize".to_string(),
1565                            json!({}),
1566                            RequestId::Number(1),
1567                        )
1568                        .unwrap(),
1569                    ))
1570                    .map_err(|e| {
1571                        AcpError::internal_error().data(format!("send initialize: {e}"))
1572                    })?;
1573                into_single_message(channel.rx.next().await.ok_or_else(|| {
1574                    AcpError::internal_error().data("initialize response channel closed")
1575                })?)?;
1576
1577                for method in ["custom/first", "custom/second"] {
1578                    channel
1579                        .tx
1580                        .unbounded_send(single_frame(
1581                            RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(),
1582                        ))
1583                        .map_err(|e| {
1584                            AcpError::internal_error().data(format!("send {method}: {e}"))
1585                        })?;
1586                }
1587
1588                finish.notified().await;
1589                finished.notify_one();
1590                Ok(())
1591            };
1592
1593            let ((), ()) = futures::try_join!(transport, client)?;
1594            Ok(())
1595        }
1596    }
1597
1598    impl ConnectTo<Agent> for InitializeThenExitClient {
1599        async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), AcpError> {
1600            let Self {
1601                sse_started,
1602                finished,
1603            } = self;
1604            let (mut channel, transport) = agent.into_channel_and_future();
1605            let client = async move {
1606                channel
1607                    .tx
1608                    .unbounded_send(single_frame(
1609                        RawJsonRpcMessage::request(
1610                            "initialize".to_string(),
1611                            json!({}),
1612                            RequestId::Number(1),
1613                        )
1614                        .unwrap(),
1615                    ))
1616                    .map_err(|error| {
1617                        AcpError::internal_error().data(format!("send initialize: {error}"))
1618                    })?;
1619                into_single_message(channel.rx.next().await.ok_or_else(|| {
1620                    AcpError::internal_error().data("initialize response channel closed")
1621                })?)?;
1622
1623                sse_started.notified().await;
1624                finished.notify_one();
1625                Ok(())
1626            };
1627
1628            let ((), ()) = futures::try_join!(transport, client)?;
1629            Ok(())
1630        }
1631    }
1632
1633    #[test]
1634    fn new_targets_standard_acp_endpoint() {
1635        assert_eq!(
1636            HttpClient::new("http://example.com")
1637                .unwrap()
1638                .endpoint
1639                .as_str(),
1640            "http://example.com/acp"
1641        );
1642        assert_eq!(
1643            HttpClient::new("http://example.com/proxy")
1644                .unwrap()
1645                .endpoint
1646                .as_str(),
1647            "http://example.com/proxy/acp"
1648        );
1649        assert_eq!(
1650            HttpClient::new("http://example.com/proxy/acp")
1651                .unwrap()
1652                .endpoint
1653                .as_str(),
1654            "http://example.com/proxy/acp"
1655        );
1656    }
1657
1658    #[test]
1659    fn with_endpoint_preserves_explicit_endpoint_path() {
1660        assert_eq!(
1661            HttpClient::with_endpoint("http://example.com/agent")
1662                .unwrap()
1663                .endpoint
1664                .as_str(),
1665            "http://example.com/agent"
1666        );
1667        assert_eq!(
1668            HttpClient::with_endpoint_and_client(
1669                "ws://example.com/custom/acp?token=abc",
1670                reqwest::Client::new(),
1671            )
1672            .unwrap()
1673            .endpoint
1674            .as_str(),
1675            "ws://example.com/custom/acp?token=abc"
1676        );
1677    }
1678
1679    #[tokio::test]
1680    async fn post_sends_cancel_request_without_session_header() {
1681        let (capture_tx, mut capture_rx) = tokio::sync::mpsc::unbounded_channel();
1682        let post_count = Arc::new(AtomicUsize::new(0));
1683        let app = Router::new().route(
1684            "/acp",
1685            post({
1686                let capture_tx = capture_tx.clone();
1687                let post_count = post_count.clone();
1688                move |headers: HeaderMap, Json(message): Json<RawJsonRpcMessage>| {
1689                    let capture_tx = capture_tx.clone();
1690                    let post_count = post_count.clone();
1691                    async move {
1692                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1693                            return initialize_response().await.into_response();
1694                        }
1695
1696                        capture_tx
1697                            .send((headers.get(HEADER_SESSION_ID).cloned(), message))
1698                            .unwrap();
1699                        StatusCode::ACCEPTED.into_response()
1700                    }
1701                }
1702            })
1703            .get(pending_sse)
1704            .delete(|| async { StatusCode::ACCEPTED }),
1705        );
1706        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1707        let addr = listener.local_addr().unwrap();
1708        let server = tokio::spawn(async move {
1709            axum::serve(listener, app).await.unwrap();
1710        });
1711        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1712        let (mut caller, transport) = Channel::duplex();
1713        let transport = tokio::spawn(run(client, transport));
1714
1715        caller
1716            .tx
1717            .unbounded_send(single_frame(
1718                RawJsonRpcMessage::request(
1719                    "initialize".to_string(),
1720                    json!({}),
1721                    RequestId::Number(1),
1722                )
1723                .unwrap(),
1724            ))
1725            .unwrap();
1726        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1727            .await
1728            .unwrap()
1729            .unwrap()
1730            .unwrap();
1731        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1732
1733        caller
1734            .tx
1735            .unbounded_send(single_frame(
1736                RawJsonRpcMessage::notification(
1737                    "$/cancel_request".to_string(),
1738                    json!({
1739                        "requestId": 2,
1740                        "sessionId": "session-1"
1741                    }),
1742                )
1743                .unwrap(),
1744            ))
1745            .unwrap();
1746
1747        let (session_header, message) = timeout(Duration::from_secs(1), capture_rx.recv())
1748            .await
1749            .unwrap()
1750            .unwrap();
1751        assert!(session_header.is_none());
1752        assert!(matches!(
1753            message,
1754            RawJsonRpcMessage::Notification(notification)
1755                if notification.method.as_ref() == "$/cancel_request"
1756        ));
1757
1758        drop(caller);
1759        timeout(Duration::from_secs(1), transport)
1760            .await
1761            .unwrap()
1762            .unwrap()
1763            .unwrap();
1764
1765        server.abort();
1766    }
1767
1768    #[tokio::test]
1769    async fn http_preserves_batch_frames_across_post_and_sse() {
1770        let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel();
1771        let post_count = Arc::new(AtomicUsize::new(0));
1772        let emit_sse = Arc::new(Notify::new());
1773        let inbound_batch = json!([
1774            {
1775                "jsonrpc": "2.0",
1776                "method": "custom/inbound-one",
1777                "params": {}
1778            },
1779            {
1780                "jsonrpc": "2.0",
1781                "method": "custom/inbound-two",
1782                "params": {}
1783            }
1784        ]);
1785        let app = Router::new().route(
1786            "/acp",
1787            post({
1788                let post_count = post_count.clone();
1789                move |body: String| {
1790                    let post_count = post_count.clone();
1791                    let post_tx = post_tx.clone();
1792                    async move {
1793                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1794                            return initialize_response().await.into_response();
1795                        }
1796
1797                        post_tx
1798                            .send(serde_json::from_str::<serde_json::Value>(&body).unwrap())
1799                            .unwrap();
1800                        StatusCode::ACCEPTED.into_response()
1801                    }
1802                }
1803            })
1804            .get({
1805                let emit_sse = emit_sse.clone();
1806                let inbound_batch = inbound_batch.clone();
1807                move || {
1808                    let emit_sse = emit_sse.clone();
1809                    let inbound_batch = inbound_batch.clone();
1810                    async move {
1811                        let stream = async_stream::stream! {
1812                            emit_sse.notified().await;
1813                            yield Ok::<_, Infallible>(
1814                                Event::default().data(inbound_batch.to_string()),
1815                            );
1816                            futures::future::pending::<()>().await;
1817                        };
1818                        Sse::new(stream)
1819                    }
1820                }
1821            })
1822            .delete(|| async { StatusCode::ACCEPTED }),
1823        );
1824        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1825        let addr = listener.local_addr().unwrap();
1826        let server = tokio::spawn(async move {
1827            axum::serve(listener, app).await.unwrap();
1828        });
1829        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1830        let (mut caller, transport) = Channel::duplex();
1831        let transport = tokio::spawn(run(client, transport));
1832
1833        caller
1834            .tx
1835            .unbounded_send(single_frame(
1836                RawJsonRpcMessage::request(
1837                    "initialize".to_string(),
1838                    json!({}),
1839                    RequestId::Number(1),
1840                )
1841                .unwrap(),
1842            ))
1843            .unwrap();
1844        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1845            .await
1846            .unwrap()
1847            .unwrap()
1848            .unwrap();
1849        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1850
1851        let outbound_batch = json!([
1852            {
1853                "jsonrpc": "2.0",
1854                "method": "custom/outbound-one",
1855                "params": {}
1856            },
1857            {
1858                "jsonrpc": "2.0",
1859                "method": "custom/outbound-two",
1860                "params": {}
1861            }
1862        ]);
1863        caller
1864            .tx
1865            .unbounded_send(TransportFrame::Batch(
1866                TransportBatch::from_messages([
1867                    RawJsonRpcMessage::notification("custom/outbound-one".to_string(), json!({}))
1868                        .unwrap(),
1869                    RawJsonRpcMessage::notification("custom/outbound-two".to_string(), json!({}))
1870                        .unwrap(),
1871                ])
1872                .unwrap(),
1873            ))
1874            .unwrap();
1875
1876        let posted = timeout(Duration::from_secs(1), post_rx.recv())
1877            .await
1878            .unwrap()
1879            .unwrap();
1880        assert_eq!(posted, outbound_batch);
1881
1882        emit_sse.notify_one();
1883        let inbound = timeout(Duration::from_secs(1), caller.rx.next())
1884            .await
1885            .unwrap()
1886            .unwrap();
1887        assert!(matches!(&inbound, TransportFrame::Batch(_)));
1888        assert_eq!(
1889            serde_json::from_str::<serde_json::Value>(&inbound.to_json().unwrap()).unwrap(),
1890            inbound_batch
1891        );
1892
1893        drop(caller);
1894        timeout(Duration::from_secs(1), transport)
1895            .await
1896            .unwrap()
1897            .unwrap()
1898            .unwrap();
1899
1900        server.abort();
1901    }
1902
1903    #[tokio::test]
1904    async fn batch_fork_opens_source_and_result_session_streams() {
1905        let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel();
1906        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
1907        let post_count = Arc::new(AtomicUsize::new(0));
1908        let emit_response = Arc::new(Notify::new());
1909        let connection_stream_established = Arc::new(AtomicBool::new(false));
1910        let source_stream_established = Arc::new(AtomicBool::new(false));
1911        let response_batch = json!([
1912            {
1913                "jsonrpc": "2.0",
1914                "id": 2,
1915                "result": { "sessionId": "forked-session" }
1916            }
1917        ]);
1918        let app = Router::new().route(
1919            "/acp",
1920            post({
1921                let post_count = post_count.clone();
1922                let connection_stream_established = connection_stream_established.clone();
1923                let source_stream_established = source_stream_established.clone();
1924                move |body: String| {
1925                    let post_count = post_count.clone();
1926                    let post_tx = post_tx.clone();
1927                    let connection_stream_established = connection_stream_established.clone();
1928                    let source_stream_established = source_stream_established.clone();
1929                    async move {
1930                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1931                            return initialize_response().await.into_response();
1932                        }
1933
1934                        if !connection_stream_established.load(Ordering::SeqCst)
1935                            || !source_stream_established.load(Ordering::SeqCst)
1936                        {
1937                            return StatusCode::CONFLICT.into_response();
1938                        }
1939                        post_tx
1940                            .send(serde_json::from_str::<serde_json::Value>(&body).unwrap())
1941                            .unwrap();
1942                        StatusCode::ACCEPTED.into_response()
1943                    }
1944                }
1945            })
1946            .get({
1947                let emit_response = emit_response.clone();
1948                let response_batch = response_batch.clone();
1949                let connection_stream_established = connection_stream_established.clone();
1950                let source_stream_established = source_stream_established.clone();
1951                move |headers: HeaderMap| {
1952                    let emit_response = emit_response.clone();
1953                    let response_batch = response_batch.clone();
1954                    let get_tx = get_tx.clone();
1955                    let connection_stream_established = connection_stream_established.clone();
1956                    let source_stream_established = source_stream_established.clone();
1957                    async move {
1958                        let session_id = headers
1959                            .get(HEADER_SESSION_ID)
1960                            .and_then(|value| value.to_str().ok())
1961                            .map(String::from);
1962                        let is_connection_stream = session_id.is_none();
1963                        let is_source_stream = session_id.as_deref() == Some("source-session");
1964                        if is_connection_stream {
1965                            sleep(Duration::from_millis(50)).await;
1966                            connection_stream_established.store(true, Ordering::SeqCst);
1967                        }
1968                        if is_source_stream {
1969                            sleep(Duration::from_millis(50)).await;
1970                            source_stream_established.store(true, Ordering::SeqCst);
1971                        }
1972                        get_tx.send(session_id).unwrap();
1973
1974                        let stream = async_stream::stream! {
1975                            if is_source_stream {
1976                                emit_response.notified().await;
1977                                yield Ok::<_, Infallible>(
1978                                    Event::default().data(response_batch.to_string()),
1979                                );
1980                            }
1981                            futures::future::pending::<()>().await;
1982                        };
1983                        Sse::new(stream)
1984                    }
1985                }
1986            })
1987            .delete(|| async { StatusCode::ACCEPTED }),
1988        );
1989        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1990        let addr = listener.local_addr().unwrap();
1991        let server = tokio::spawn(async move {
1992            axum::serve(listener, app).await.unwrap();
1993        });
1994        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1995        let (mut caller, transport) = Channel::duplex();
1996        let transport = tokio::spawn(run(client, transport));
1997
1998        caller
1999            .tx
2000            .unbounded_send(single_frame(
2001                RawJsonRpcMessage::request(
2002                    "initialize".to_string(),
2003                    json!({}),
2004                    RequestId::Number(1),
2005                )
2006                .unwrap(),
2007            ))
2008            .unwrap();
2009        timeout(Duration::from_secs(1), caller.rx.next())
2010            .await
2011            .unwrap()
2012            .unwrap();
2013
2014        caller
2015            .tx
2016            .unbounded_send(TransportFrame::Batch(
2017                TransportBatch::from_messages([RawJsonRpcMessage::request(
2018                    "session/fork".to_string(),
2019                    json!({ "sessionId": "source-session" }),
2020                    RequestId::Number(2),
2021                )
2022                .unwrap()])
2023                .unwrap(),
2024            ))
2025            .unwrap();
2026
2027        let connection_stream = timeout(Duration::from_secs(1), get_rx.recv())
2028            .await
2029            .unwrap()
2030            .unwrap();
2031        assert!(connection_stream.is_none());
2032        let source_stream = timeout(Duration::from_secs(1), get_rx.recv())
2033            .await
2034            .unwrap()
2035            .unwrap();
2036        assert_eq!(source_stream.as_deref(), Some("source-session"));
2037        let posted = timeout(Duration::from_secs(1), post_rx.recv())
2038            .await
2039            .unwrap()
2040            .unwrap();
2041        assert!(posted.is_array(), "outgoing batch must remain an array");
2042
2043        emit_response.notify_one();
2044        let response = timeout(Duration::from_secs(1), caller.rx.next())
2045            .await
2046            .unwrap()
2047            .unwrap();
2048        assert!(matches!(&response, TransportFrame::Batch(_)));
2049        assert_eq!(
2050            serde_json::from_str::<serde_json::Value>(&response.to_json().unwrap()).unwrap(),
2051            response_batch
2052        );
2053        let forked_stream = timeout(Duration::from_secs(1), get_rx.recv())
2054            .await
2055            .unwrap()
2056            .unwrap();
2057        assert_eq!(forked_stream.as_deref(), Some("forked-session"));
2058
2059        drop(caller);
2060        timeout(Duration::from_secs(1), transport)
2061            .await
2062            .unwrap()
2063            .unwrap()
2064            .unwrap();
2065
2066        server.abort();
2067    }
2068
2069    #[tokio::test]
2070    async fn custom_response_with_session_id_does_not_open_session_sse() {
2071        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
2072        let response_ready = Arc::new(tokio::sync::Notify::new());
2073        let post_count = Arc::new(AtomicUsize::new(0));
2074        let app = Router::new().route(
2075            "/acp",
2076            post({
2077                let post_count = post_count.clone();
2078                let response_ready = response_ready.clone();
2079                move |Json(_message): Json<RawJsonRpcMessage>| {
2080                    let post_count = post_count.clone();
2081                    let response_ready = response_ready.clone();
2082                    async move {
2083                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2084                            return initialize_response().await.into_response();
2085                        }
2086
2087                        response_ready.notify_waiters();
2088                        StatusCode::ACCEPTED.into_response()
2089                    }
2090                }
2091            })
2092            .get({
2093                let get_tx = get_tx.clone();
2094                let response_ready = response_ready.clone();
2095                move |headers: HeaderMap| {
2096                    let get_tx = get_tx.clone();
2097                    let response_ready = response_ready.clone();
2098                    async move {
2099                        let session_header = headers
2100                            .get(HEADER_SESSION_ID)
2101                            .and_then(|value| value.to_str().ok())
2102                            .map(String::from);
2103                        get_tx.send(session_header).unwrap();
2104
2105                        let stream = async_stream::stream! {
2106                            response_ready.notified().await;
2107                            yield Ok::<_, Infallible>(sse_event(
2108                                RawJsonRpcMessage::response(
2109                                    RequestId::Number(2),
2110                                    Ok(json!({ "sessionId": "session-1" })),
2111                                ),
2112                            ));
2113                            futures::future::pending::<()>().await;
2114                        };
2115                        Sse::new(stream)
2116                    }
2117                }
2118            })
2119            .delete(|| async { StatusCode::ACCEPTED }),
2120        );
2121        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2122        let addr = listener.local_addr().unwrap();
2123        let server = tokio::spawn(async move {
2124            axum::serve(listener, app).await.unwrap();
2125        });
2126        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2127        let (mut caller, transport) = Channel::duplex();
2128        let transport = tokio::spawn(run(client, transport));
2129
2130        caller
2131            .tx
2132            .unbounded_send(single_frame(
2133                RawJsonRpcMessage::request(
2134                    "initialize".to_string(),
2135                    json!({}),
2136                    RequestId::Number(1),
2137                )
2138                .unwrap(),
2139            ))
2140            .unwrap();
2141        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2142            .await
2143            .unwrap()
2144            .unwrap()
2145            .unwrap();
2146        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2147
2148        let connection_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2149            .await
2150            .unwrap()
2151            .unwrap();
2152        assert!(connection_sse_header.is_none());
2153
2154        caller
2155            .tx
2156            .unbounded_send(single_frame(
2157                RawJsonRpcMessage::request(
2158                    "custom/sessionish".to_string(),
2159                    json!({}),
2160                    RequestId::Number(2),
2161                )
2162                .unwrap(),
2163            ))
2164            .unwrap();
2165        let response = timeout(Duration::from_secs(1), caller.rx.next())
2166            .await
2167            .unwrap()
2168            .unwrap()
2169            .unwrap();
2170        assert!(matches!(
2171            response,
2172            RawJsonRpcMessage::Response(RpcResponse::Result {
2173                id: RequestId::Number(2),
2174                ..
2175            })
2176        ));
2177
2178        assert!(
2179            timeout(Duration::from_millis(100), get_rx.recv())
2180                .await
2181                .is_err(),
2182            "custom response must not open a session SSE stream"
2183        );
2184
2185        drop(caller);
2186        timeout(Duration::from_secs(1), transport)
2187            .await
2188            .unwrap()
2189            .unwrap()
2190            .unwrap();
2191
2192        server.abort();
2193    }
2194
2195    #[tokio::test]
2196    async fn fork_response_with_session_id_opens_session_sse() {
2197        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
2198        let response_ready = Arc::new(tokio::sync::Notify::new());
2199        let post_count = Arc::new(AtomicUsize::new(0));
2200        let app = Router::new().route(
2201            "/acp",
2202            post({
2203                let post_count = post_count.clone();
2204                let response_ready = response_ready.clone();
2205                move |Json(_message): Json<RawJsonRpcMessage>| {
2206                    let post_count = post_count.clone();
2207                    let response_ready = response_ready.clone();
2208                    async move {
2209                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2210                            return initialize_response().await.into_response();
2211                        }
2212
2213                        response_ready.notify_waiters();
2214                        StatusCode::ACCEPTED.into_response()
2215                    }
2216                }
2217            })
2218            .get({
2219                let get_tx = get_tx.clone();
2220                let response_ready = response_ready.clone();
2221                move |headers: HeaderMap| {
2222                    let get_tx = get_tx.clone();
2223                    let response_ready = response_ready.clone();
2224                    async move {
2225                        let session_header = headers
2226                            .get(HEADER_SESSION_ID)
2227                            .and_then(|value| value.to_str().ok())
2228                            .map(String::from);
2229                        let is_connection_stream = session_header.is_none();
2230                        get_tx.send(session_header).unwrap();
2231
2232                        let stream = async_stream::stream! {
2233                            if is_connection_stream {
2234                                response_ready.notified().await;
2235                                yield Ok::<_, Infallible>(sse_event(
2236                                    RawJsonRpcMessage::response(
2237                                        RequestId::Number(2),
2238                                        Ok(json!({ "sessionId": "forked-session" })),
2239                                    ),
2240                                ));
2241                            }
2242                            futures::future::pending::<()>().await;
2243                        };
2244                        Sse::new(stream)
2245                    }
2246                }
2247            })
2248            .delete(|| async { StatusCode::ACCEPTED }),
2249        );
2250        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2251        let addr = listener.local_addr().unwrap();
2252        let server = tokio::spawn(async move {
2253            axum::serve(listener, app).await.unwrap();
2254        });
2255        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2256        let (mut caller, transport) = Channel::duplex();
2257        let transport = tokio::spawn(run(client, transport));
2258
2259        caller
2260            .tx
2261            .unbounded_send(single_frame(
2262                RawJsonRpcMessage::request(
2263                    "initialize".to_string(),
2264                    json!({}),
2265                    RequestId::Number(1),
2266                )
2267                .unwrap(),
2268            ))
2269            .unwrap();
2270        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2271            .await
2272            .unwrap()
2273            .unwrap()
2274            .unwrap();
2275        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2276
2277        let connection_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2278            .await
2279            .unwrap()
2280            .unwrap();
2281        assert!(connection_sse_header.is_none());
2282
2283        caller
2284            .tx
2285            .unbounded_send(single_frame(
2286                RawJsonRpcMessage::request(
2287                    "session/fork".to_string(),
2288                    json!({ "sessionId": "source-session" }),
2289                    RequestId::Number(2),
2290                )
2291                .unwrap(),
2292            ))
2293            .unwrap();
2294        let source_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2295            .await
2296            .unwrap()
2297            .unwrap();
2298        assert_eq!(source_sse_header.as_deref(), Some("source-session"));
2299
2300        let response = timeout(Duration::from_secs(1), caller.rx.next())
2301            .await
2302            .unwrap()
2303            .unwrap()
2304            .unwrap();
2305        assert!(matches!(
2306            response,
2307            RawJsonRpcMessage::Response(RpcResponse::Result {
2308                id: RequestId::Number(2),
2309                ..
2310            })
2311        ));
2312
2313        let fork_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2314            .await
2315            .unwrap()
2316            .unwrap();
2317        assert_eq!(fork_sse_header.as_deref(), Some("forked-session"));
2318
2319        drop(caller);
2320        timeout(Duration::from_secs(1), transport)
2321            .await
2322            .unwrap()
2323            .unwrap()
2324            .unwrap();
2325
2326        server.abort();
2327    }
2328
2329    #[tokio::test]
2330    async fn only_response_batches_bypass_ordered_posts() {
2331        let slow_started = Arc::new(Notify::new());
2332        let release_slow = Arc::new(Notify::new());
2333        let call_batch_seen = Arc::new(Notify::new());
2334        let response_batch_seen = Arc::new(Notify::new());
2335        let app = Router::new().route(
2336            "/acp",
2337            post({
2338                let slow_started = slow_started.clone();
2339                let release_slow = release_slow.clone();
2340                let call_batch_seen = call_batch_seen.clone();
2341                let response_batch_seen = response_batch_seen.clone();
2342                move |body: String| {
2343                    let slow_started = slow_started.clone();
2344                    let release_slow = release_slow.clone();
2345                    let call_batch_seen = call_batch_seen.clone();
2346                    let response_batch_seen = response_batch_seen.clone();
2347                    async move {
2348                        let value = serde_json::from_str::<serde_json::Value>(&body).unwrap();
2349                        if value.get("method").and_then(serde_json::Value::as_str)
2350                            == Some("initialize")
2351                        {
2352                            return initialize_response().await.into_response();
2353                        }
2354                        if value.get("method").and_then(serde_json::Value::as_str)
2355                            == Some("custom/slow")
2356                        {
2357                            slow_started.notify_one();
2358                            release_slow.notified().await;
2359                        } else if let Some(entries) = value.as_array() {
2360                            if entries.iter().all(|entry| entry.get("method").is_none()) {
2361                                response_batch_seen.notify_one();
2362                            } else {
2363                                call_batch_seen.notify_one();
2364                            }
2365                        }
2366                        StatusCode::ACCEPTED.into_response()
2367                    }
2368                }
2369            })
2370            .get(pending_sse)
2371            .delete(|| async { StatusCode::ACCEPTED }),
2372        );
2373        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2374        let addr = listener.local_addr().unwrap();
2375        let server = tokio::spawn(async move {
2376            axum::serve(listener, app).await.unwrap();
2377        });
2378        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2379        let (mut caller, transport) = Channel::duplex();
2380        let transport = tokio::spawn(run(client, transport));
2381
2382        caller
2383            .tx
2384            .unbounded_send(single_frame(
2385                RawJsonRpcMessage::request(
2386                    "initialize".to_string(),
2387                    json!({}),
2388                    RequestId::Number(1),
2389                )
2390                .unwrap(),
2391            ))
2392            .unwrap();
2393        timeout(Duration::from_secs(1), caller.rx.next())
2394            .await
2395            .unwrap()
2396            .unwrap();
2397
2398        caller
2399            .tx
2400            .unbounded_send(single_frame(
2401                RawJsonRpcMessage::notification("custom/slow".to_string(), json!({})).unwrap(),
2402            ))
2403            .unwrap();
2404        timeout(Duration::from_secs(1), slow_started.notified())
2405            .await
2406            .unwrap();
2407
2408        caller
2409            .tx
2410            .unbounded_send(TransportFrame::Batch(
2411                TransportBatch::from_messages([
2412                    RawJsonRpcMessage::notification("custom/one".to_string(), json!({})).unwrap(),
2413                    RawJsonRpcMessage::notification("custom/two".to_string(), json!({})).unwrap(),
2414                ])
2415                .unwrap(),
2416            ))
2417            .unwrap();
2418        assert!(
2419            timeout(Duration::from_millis(100), call_batch_seen.notified())
2420                .await
2421                .is_err(),
2422            "call-bearing batches must remain behind an earlier ordered POST"
2423        );
2424
2425        caller
2426            .tx
2427            .unbounded_send(TransportFrame::Batch(
2428                TransportBatch::from_messages([
2429                    RawJsonRpcMessage::response(RequestId::Number(10), Ok(json!({}))),
2430                    RawJsonRpcMessage::response(RequestId::Number(11), Ok(json!({}))),
2431                ])
2432                .unwrap(),
2433            ))
2434            .unwrap();
2435        timeout(Duration::from_secs(1), response_batch_seen.notified())
2436            .await
2437            .expect("response-only batch should bypass the ordered POST queue");
2438
2439        release_slow.notify_one();
2440        timeout(Duration::from_secs(1), call_batch_seen.notified())
2441            .await
2442            .expect("call-bearing batch should be sent after the earlier POST completes");
2443
2444        drop(caller);
2445        timeout(Duration::from_secs(1), transport)
2446            .await
2447            .unwrap()
2448            .unwrap()
2449            .unwrap();
2450
2451        server.abort();
2452    }
2453
2454    #[tokio::test]
2455    async fn client_completion_drains_ordered_posts_in_order() {
2456        let first_started = Arc::new(Notify::new());
2457        let release_first = Arc::new(Notify::new());
2458        let second_seen = Arc::new(Notify::new());
2459        let finish_client = Arc::new(Notify::new());
2460        let client_finished = Arc::new(Notify::new());
2461        let (escaped_tx, escaped_rx) = futures::channel::oneshot::channel();
2462        let app = Router::new().route(
2463            "/acp",
2464            post({
2465                let first_started = first_started.clone();
2466                let release_first = release_first.clone();
2467                let second_seen = second_seen.clone();
2468                move |Json(message): Json<RawJsonRpcMessage>| {
2469                    let first_started = first_started.clone();
2470                    let release_first = release_first.clone();
2471                    let second_seen = second_seen.clone();
2472                    async move {
2473                        if is_initialize_request(&message) {
2474                            return initialize_response().await.into_response();
2475                        }
2476
2477                        match method_for_message(&message) {
2478                            Some("custom/first") => {
2479                                first_started.notify_one();
2480                                release_first.notified().await;
2481                            }
2482                            Some("custom/second") => {
2483                                second_seen.notify_one();
2484                            }
2485                            _ => {}
2486                        }
2487                        StatusCode::ACCEPTED.into_response()
2488                    }
2489                }
2490            })
2491            .get(pending_sse)
2492            .delete(|| async { StatusCode::ACCEPTED }),
2493        );
2494        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2495        let addr = listener.local_addr().unwrap();
2496        let server = tokio::spawn(async move {
2497            axum::serve(listener, app).await.unwrap();
2498        });
2499        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2500        let mut connection = tokio::spawn(client.connect_to(PostsThenExitClient {
2501            finish: finish_client.clone(),
2502            finished: client_finished.clone(),
2503            escaped_tx,
2504        }));
2505        let escaped = timeout(Duration::from_secs(1), escaped_rx)
2506            .await
2507            .unwrap()
2508            .unwrap();
2509
2510        timeout(Duration::from_secs(1), first_started.notified())
2511            .await
2512            .unwrap();
2513        assert!(
2514            timeout(Duration::from_millis(100), second_seen.notified())
2515                .await
2516                .is_err(),
2517            "second POST must not be sent while the first POST is pending"
2518        );
2519
2520        finish_client.notify_one();
2521        timeout(Duration::from_secs(1), client_finished.notified())
2522            .await
2523            .unwrap();
2524        assert!(
2525            timeout(Duration::from_millis(100), &mut connection)
2526                .await
2527                .is_err(),
2528            "HTTP transport returned before its accepted POSTs completed"
2529        );
2530        assert!(
2531            escaped
2532                .unbounded_send(single_frame(
2533                    RawJsonRpcMessage::notification("custom/too-late".to_string(), json!({}),)
2534                        .unwrap()
2535                ))
2536                .is_err(),
2537            "escaped client sender remained open after client completion"
2538        );
2539
2540        release_first.notify_one();
2541        timeout(Duration::from_secs(1), second_seen.notified())
2542            .await
2543            .unwrap();
2544
2545        timeout(Duration::from_secs(1), connection)
2546            .await
2547            .unwrap()
2548            .unwrap()
2549            .unwrap();
2550
2551        server.abort();
2552    }
2553
2554    #[tokio::test]
2555    async fn client_completion_cancels_pending_sse_establishment() {
2556        let sse_started = Arc::new(Notify::new());
2557        let delete_count = Arc::new(AtomicUsize::new(0));
2558        let client_finished = Arc::new(Notify::new());
2559        let app = Router::new().route(
2560            "/acp",
2561            post(initialize_response)
2562                .get({
2563                    let sse_started = sse_started.clone();
2564                    move || {
2565                        let sse_started = sse_started.clone();
2566                        async move {
2567                            sse_started.notify_one();
2568                            futures::future::pending::<StatusCode>().await
2569                        }
2570                    }
2571                })
2572                .delete({
2573                    let delete_count = delete_count.clone();
2574                    move || {
2575                        let delete_count = delete_count.clone();
2576                        async move {
2577                            delete_count.fetch_add(1, Ordering::SeqCst);
2578                            StatusCode::ACCEPTED
2579                        }
2580                    }
2581                }),
2582        );
2583        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2584        let addr = listener.local_addr().unwrap();
2585        let server = tokio::spawn(async move {
2586            axum::serve(listener, app).await.unwrap();
2587        });
2588        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2589        let connection = tokio::spawn(client.connect_to(InitializeThenExitClient {
2590            sse_started,
2591            finished: client_finished.clone(),
2592        }));
2593
2594        timeout(Duration::from_secs(1), client_finished.notified())
2595            .await
2596            .expect("client foreground did not finish after the SSE request started");
2597
2598        timeout(Duration::from_secs(1), connection)
2599            .await
2600            .expect("transport remained blocked on SSE response headers")
2601            .unwrap()
2602            .unwrap();
2603        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
2604
2605        server.abort();
2606    }
2607
2608    #[tokio::test]
2609    async fn stalled_sse_establishment_observes_earlier_post_failure() {
2610        let app = Router::new().route(
2611            "/acp",
2612            get(|| async { futures::future::pending::<StatusCode>().await }),
2613        );
2614        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2615        let addr = listener.local_addr().unwrap();
2616        let server = tokio::spawn(async move {
2617            axum::serve(listener, app).await.unwrap();
2618        });
2619
2620        let connection = HttpConnection::new(
2621            url::Url::parse(&format!("http://{addr}/acp")).unwrap(),
2622            reqwest::Client::new(),
2623        );
2624        connection.set_connection_id("connection-1".to_string());
2625        let (incoming, _incoming_rx) = mpsc::unbounded();
2626        let mut state = ClientState {
2627            connection: connection.clone(),
2628            open_session_streams: HashSet::new(),
2629            pending_requests: HashMap::new(),
2630            incoming,
2631        };
2632        let pending_request = (RequestId::Number(7), "custom/earlier".to_string());
2633        state.track_pending_requests(std::slice::from_ref(&pending_request));
2634        let mut posts = PostQueues::default();
2635        posts.ordered.push(PendingPost {
2636            pending_requests: vec![pending_request],
2637            response: async { Err("earlier post failed".to_string()) }.boxed(),
2638        });
2639
2640        let (_outgoing_tx, mut outgoing) = mpsc::unbounded();
2641        let mut buffered_outgoing = VecDeque::new();
2642        let (event_tx, mut event_rx) = mpsc::unbounded();
2643        let mut lifecycle = HttpTransportLifecycle::new(connection);
2644        let error = timeout(
2645            Duration::from_secs(1),
2646            lifecycle.start_sse(
2647                Some("later-session".to_string()),
2648                event_tx,
2649                SseStartContext {
2650                    events: &mut event_rx,
2651                    outgoing: &mut outgoing,
2652                    buffered_outgoing: &mut buffered_outgoing,
2653                    posts: &mut posts,
2654                    state: &mut state,
2655                },
2656            ),
2657        )
2658        .await
2659        .expect("stalled SSE setup hid an earlier POST failure")
2660        .unwrap_err();
2661
2662        assert!(error.to_string().contains("earlier post failed"));
2663        assert!(state.pending_requests.is_empty());
2664
2665        lifecycle.close().await;
2666        server.abort();
2667    }
2668
2669    #[tokio::test]
2670    async fn stalled_sse_establishment_keeps_callback_responses_moving() {
2671        let release_get = Arc::new(Notify::new());
2672        let complete_earlier_post = Arc::new(Notify::new());
2673        let app = Router::new().route(
2674            "/acp",
2675            post({
2676                let release_get = release_get.clone();
2677                let complete_earlier_post = complete_earlier_post.clone();
2678                move || {
2679                    release_get.notify_one();
2680                    complete_earlier_post.notify_one();
2681                    async { StatusCode::ACCEPTED }
2682                }
2683            })
2684            .get({
2685                let release_get = release_get.clone();
2686                move || {
2687                    let release_get = release_get.clone();
2688                    async move {
2689                        release_get.notified().await;
2690                        Sse::new(futures::stream::pending::<Result<Event, Infallible>>())
2691                    }
2692                }
2693            }),
2694        );
2695        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2696        let addr = listener.local_addr().unwrap();
2697        let server = tokio::spawn(async move {
2698            axum::serve(listener, app).await.unwrap();
2699        });
2700
2701        let connection = HttpConnection::new(
2702            url::Url::parse(&format!("http://{addr}/acp")).unwrap(),
2703            reqwest::Client::new(),
2704        );
2705        connection.set_connection_id("connection-1".to_string());
2706        let (incoming, mut incoming_rx) = mpsc::unbounded();
2707        let mut state = ClientState {
2708            connection: connection.clone(),
2709            open_session_streams: HashSet::new(),
2710            pending_requests: HashMap::new(),
2711            incoming,
2712        };
2713        let mut posts = PostQueues::default();
2714        posts.ordered.push(PendingPost {
2715            pending_requests: Vec::new(),
2716            response: async move {
2717                complete_earlier_post.notified().await;
2718                Ok(())
2719            }
2720            .boxed(),
2721        });
2722
2723        let (outgoing_tx, mut outgoing) = mpsc::unbounded();
2724        let outgoing_guard = outgoing_tx.clone();
2725        let mut buffered_outgoing = VecDeque::new();
2726        let (event_tx, mut event_rx) = mpsc::unbounded();
2727        event_tx
2728            .unbounded_send(SseMessage {
2729                frame: single_frame(
2730                    RawJsonRpcMessage::request(
2731                        "test/callback".to_string(),
2732                        json!({}),
2733                        RequestId::Number(99),
2734                    )
2735                    .unwrap(),
2736                ),
2737            })
2738            .unwrap();
2739
2740        let responder = async move {
2741            let callback = incoming_rx
2742                .next()
2743                .await
2744                .expect("callback was not delivered");
2745            assert!(matches!(
2746                into_single_message(callback).unwrap(),
2747                RawJsonRpcMessage::Request(request)
2748                    if request.method.as_ref() == "test/callback"
2749            ));
2750            outgoing_tx
2751                .unbounded_send(single_frame(RawJsonRpcMessage::response(
2752                    RequestId::Number(99),
2753                    Ok(json!({})),
2754                )))
2755                .unwrap();
2756        };
2757        let mut lifecycle = HttpTransportLifecycle::new(connection);
2758        let (outcome, ()) = timeout(Duration::from_secs(1), async {
2759            futures::join!(
2760                lifecycle.start_sse(
2761                    Some("later-session".to_string()),
2762                    event_tx,
2763                    SseStartContext {
2764                        events: &mut event_rx,
2765                        outgoing: &mut outgoing,
2766                        buffered_outgoing: &mut buffered_outgoing,
2767                        posts: &mut posts,
2768                        state: &mut state,
2769                    },
2770                ),
2771                responder,
2772            )
2773        })
2774        .await
2775        .expect("callback response deadlocked behind stalled SSE establishment");
2776
2777        assert_eq!(outcome.unwrap(), SseStartOutcome::Established);
2778        assert!(buffered_outgoing.is_empty());
2779
2780        drop(outgoing_guard);
2781        lifecycle.close().await;
2782        server.abort();
2783    }
2784
2785    #[tokio::test]
2786    async fn pending_sse_establishment_reports_buffered_output_on_shutdown() {
2787        let sse_started = Arc::new(Notify::new());
2788        let delete_count = Arc::new(AtomicUsize::new(0));
2789        let app = Router::new().route(
2790            "/acp",
2791            post(initialize_response)
2792                .get({
2793                    let sse_started = sse_started.clone();
2794                    move || {
2795                        let sse_started = sse_started.clone();
2796                        async move {
2797                            sse_started.notify_one();
2798                            futures::future::pending::<StatusCode>().await
2799                        }
2800                    }
2801                })
2802                .delete({
2803                    let delete_count = delete_count.clone();
2804                    move || {
2805                        let delete_count = delete_count.clone();
2806                        async move {
2807                            delete_count.fetch_add(1, Ordering::SeqCst);
2808                            StatusCode::ACCEPTED
2809                        }
2810                    }
2811                }),
2812        );
2813        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2814        let addr = listener.local_addr().unwrap();
2815        let server = tokio::spawn(async move {
2816            axum::serve(listener, app).await.unwrap();
2817        });
2818        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2819        let (mut caller, transport) = Channel::duplex();
2820        let transport = tokio::spawn(run(client, transport));
2821
2822        caller
2823            .tx
2824            .unbounded_send(single_frame(
2825                RawJsonRpcMessage::request(
2826                    "initialize".to_string(),
2827                    json!({}),
2828                    RequestId::Number(1),
2829                )
2830                .unwrap(),
2831            ))
2832            .unwrap();
2833        timeout(Duration::from_secs(1), caller.rx.next())
2834            .await
2835            .unwrap()
2836            .unwrap();
2837        timeout(Duration::from_secs(1), sse_started.notified())
2838            .await
2839            .expect("connection SSE request did not reach the server");
2840
2841        caller
2842            .tx
2843            .unbounded_send(single_frame(
2844                RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(),
2845            ))
2846            .unwrap();
2847        drop(caller);
2848
2849        let error = timeout(Duration::from_secs(1), transport)
2850            .await
2851            .expect("transport remained blocked on SSE response headers")
2852            .unwrap()
2853            .unwrap_err();
2854        assert!(error.to_string().contains("accepted messages"));
2855        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
2856
2857        server.abort();
2858    }
2859
2860    #[tokio::test]
2861    async fn sse_continues_while_post_is_pending() {
2862        let post_started = Arc::new(Notify::new());
2863        let callback_response_seen = Arc::new(Notify::new());
2864        let sse_started = Arc::new(Notify::new());
2865        let (callback_tx, mut callback_rx) = tokio::sync::mpsc::unbounded_channel();
2866        let app = Router::new().route(
2867            "/acp",
2868            post({
2869                let post_started = post_started.clone();
2870                let callback_response_seen = callback_response_seen.clone();
2871                let callback_tx = callback_tx.clone();
2872                move |Json(message): Json<RawJsonRpcMessage>| {
2873                    let post_started = post_started.clone();
2874                    let callback_response_seen = callback_response_seen.clone();
2875                    let callback_tx = callback_tx.clone();
2876                    async move {
2877                        if is_initialize_request(&message) {
2878                            return initialize_response().await.into_response();
2879                        }
2880
2881                        match &message {
2882                            RawJsonRpcMessage::Request(request)
2883                                if request.method.as_ref() == "custom/slow" =>
2884                            {
2885                                post_started.notify_waiters();
2886                                callback_response_seen.notified().await;
2887                                StatusCode::ACCEPTED.into_response()
2888                            }
2889                            RawJsonRpcMessage::Response(
2890                                RpcResponse::Result {
2891                                    id: RequestId::Number(99),
2892                                    ..
2893                                }
2894                                | RpcResponse::Error {
2895                                    id: RequestId::Number(99),
2896                                    ..
2897                                },
2898                            ) => {
2899                                callback_tx.send(message).unwrap();
2900                                callback_response_seen.notify_waiters();
2901                                StatusCode::ACCEPTED.into_response()
2902                            }
2903                            _ => StatusCode::ACCEPTED.into_response(),
2904                        }
2905                    }
2906                }
2907            })
2908            .get({
2909                let post_started = post_started.clone();
2910                let sse_started = sse_started.clone();
2911                move || {
2912                    let post_started = post_started.clone();
2913                    let sse_started = sse_started.clone();
2914                    async move {
2915                        let stream = async_stream::stream! {
2916                            sse_started.notify_waiters();
2917                            post_started.notified().await;
2918                            yield Ok::<_, Infallible>(sse_event(
2919                                RawJsonRpcMessage::request(
2920                                    "client/callback".to_string(),
2921                                    json!({}),
2922                                    RequestId::Number(99),
2923                                )
2924                                .unwrap(),
2925                            ));
2926                            futures::future::pending::<()>().await;
2927                        };
2928                        Sse::new(stream)
2929                    }
2930                }
2931            })
2932            .delete(|| async { StatusCode::ACCEPTED }),
2933        );
2934        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2935        let addr = listener.local_addr().unwrap();
2936        let server = tokio::spawn(async move {
2937            axum::serve(listener, app).await.unwrap();
2938        });
2939        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2940        let (mut caller, transport) = Channel::duplex();
2941        let transport = tokio::spawn(run(client, transport));
2942
2943        caller
2944            .tx
2945            .unbounded_send(single_frame(
2946                RawJsonRpcMessage::request(
2947                    "initialize".to_string(),
2948                    json!({}),
2949                    RequestId::Number(1),
2950                )
2951                .unwrap(),
2952            ))
2953            .unwrap();
2954        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2955            .await
2956            .unwrap()
2957            .unwrap()
2958            .unwrap();
2959        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2960        timeout(Duration::from_secs(1), sse_started.notified())
2961            .await
2962            .unwrap();
2963
2964        caller
2965            .tx
2966            .unbounded_send(single_frame(
2967                RawJsonRpcMessage::request(
2968                    "custom/slow".to_string(),
2969                    json!({}),
2970                    RequestId::Number(2),
2971                )
2972                .unwrap(),
2973            ))
2974            .unwrap();
2975
2976        let callback = timeout(Duration::from_secs(1), caller.rx.next())
2977            .await
2978            .unwrap()
2979            .unwrap()
2980            .unwrap();
2981        assert!(matches!(
2982            callback,
2983            RawJsonRpcMessage::Request(request)
2984                if request.method.as_ref() == "client/callback"
2985                    && request.id == RequestId::Number(99)
2986        ));
2987
2988        caller
2989            .tx
2990            .unbounded_send(single_frame(RawJsonRpcMessage::response(
2991                RequestId::Number(99),
2992                Ok(json!({})),
2993            )))
2994            .unwrap();
2995        let callback_response = timeout(Duration::from_secs(1), callback_rx.recv())
2996            .await
2997            .unwrap()
2998            .unwrap();
2999        assert!(matches!(
3000            callback_response,
3001            RawJsonRpcMessage::Response(RpcResponse::Result {
3002                id: RequestId::Number(99),
3003                ..
3004            })
3005        ));
3006
3007        drop(caller);
3008        timeout(Duration::from_secs(1), transport)
3009            .await
3010            .unwrap()
3011            .unwrap()
3012            .unwrap();
3013
3014        server.abort();
3015    }
3016
3017    #[tokio::test]
3018    async fn post_error_deletes_initialized_connection() {
3019        let delete_count = Arc::new(AtomicUsize::new(0));
3020        let delete_count_for_handler = delete_count.clone();
3021        let app = Router::new().route(
3022            "/acp",
3023            post(initialize_response).get(pending_sse).delete(move || {
3024                let delete_count = delete_count_for_handler.clone();
3025                async move {
3026                    delete_count.fetch_add(1, Ordering::SeqCst);
3027                    StatusCode::ACCEPTED
3028                }
3029            }),
3030        );
3031        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3032        let addr = listener.local_addr().unwrap();
3033        let server = tokio::spawn(async move {
3034            axum::serve(listener, app).await.unwrap();
3035        });
3036        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3037        let (mut caller, transport) = Channel::duplex();
3038        let transport = tokio::spawn(run(client, transport));
3039
3040        caller
3041            .tx
3042            .unbounded_send(single_frame(
3043                RawJsonRpcMessage::request(
3044                    "initialize".to_string(),
3045                    json!({}),
3046                    RequestId::Number(1),
3047                )
3048                .unwrap(),
3049            ))
3050            .unwrap();
3051        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3052            .await
3053            .unwrap()
3054            .unwrap()
3055            .unwrap();
3056        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3057
3058        caller
3059            .tx
3060            .unbounded_send(single_frame(
3061                RawJsonRpcMessage::request(
3062                    "session/prompt".to_string(),
3063                    json!({}),
3064                    RequestId::Number(2),
3065                )
3066                .unwrap(),
3067            ))
3068            .unwrap();
3069        let error = timeout(Duration::from_secs(1), transport)
3070            .await
3071            .unwrap()
3072            .unwrap()
3073            .unwrap_err();
3074
3075        assert!(error.to_string().contains("POST"));
3076        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3077
3078        server.abort();
3079    }
3080
3081    #[tokio::test]
3082    async fn connection_sse_disconnect_fails_transport() {
3083        let delete_count = Arc::new(AtomicUsize::new(0));
3084        let delete_count_for_handler = delete_count.clone();
3085        let app = Router::new().route(
3086            "/acp",
3087            post(initialize_response).get(closed_sse).delete(move || {
3088                let delete_count = delete_count_for_handler.clone();
3089                async move {
3090                    delete_count.fetch_add(1, Ordering::SeqCst);
3091                    StatusCode::ACCEPTED
3092                }
3093            }),
3094        );
3095        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3096        let addr = listener.local_addr().unwrap();
3097        let server = tokio::spawn(async move {
3098            axum::serve(listener, app).await.unwrap();
3099        });
3100        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3101        let (mut caller, transport) = Channel::duplex();
3102        let transport = tokio::spawn(run(client, transport));
3103
3104        caller
3105            .tx
3106            .unbounded_send(single_frame(
3107                RawJsonRpcMessage::request(
3108                    "initialize".to_string(),
3109                    json!({}),
3110                    RequestId::Number(1),
3111                )
3112                .unwrap(),
3113            ))
3114            .unwrap();
3115        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3116            .await
3117            .unwrap()
3118            .unwrap()
3119            .unwrap();
3120        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3121
3122        let error = timeout(Duration::from_secs(1), transport)
3123            .await
3124            .unwrap()
3125            .unwrap()
3126            .unwrap_err();
3127
3128        assert!(error.to_string().contains("SSE"));
3129        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3130
3131        server.abort();
3132    }
3133
3134    #[tokio::test]
3135    async fn malformed_sse_json_is_delivered_and_transport_continues() {
3136        let delete_count = Arc::new(AtomicUsize::new(0));
3137        let delete_count_for_handler = delete_count.clone();
3138        let app = Router::new().route(
3139            "/acp",
3140            post(initialize_response)
3141                .get(malformed_sse)
3142                .delete(move || {
3143                    let delete_count = delete_count_for_handler.clone();
3144                    async move {
3145                        delete_count.fetch_add(1, Ordering::SeqCst);
3146                        StatusCode::ACCEPTED
3147                    }
3148                }),
3149        );
3150        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3151        let addr = listener.local_addr().unwrap();
3152        let server = tokio::spawn(async move {
3153            axum::serve(listener, app).await.unwrap();
3154        });
3155        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3156        let (mut caller, transport) = Channel::duplex();
3157        let transport = tokio::spawn(run(client, transport));
3158
3159        caller
3160            .tx
3161            .unbounded_send(single_frame(
3162                RawJsonRpcMessage::request(
3163                    "initialize".to_string(),
3164                    json!({}),
3165                    RequestId::Number(1),
3166                )
3167                .unwrap(),
3168            ))
3169            .unwrap();
3170        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3171            .await
3172            .unwrap()
3173            .unwrap()
3174            .unwrap();
3175        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3176
3177        let frame = timeout(Duration::from_secs(1), caller.rx.next())
3178            .await
3179            .unwrap()
3180            .unwrap();
3181
3182        let TransportFrame::Malformed { raw, error } = frame else {
3183            panic!("expected malformed frame, got {frame:?}");
3184        };
3185        assert_eq!(raw, "{not json");
3186        assert_eq!(error.code, AcpError::parse_error().code);
3187        drop(caller);
3188        timeout(Duration::from_secs(1), transport)
3189            .await
3190            .unwrap()
3191            .unwrap()
3192            .unwrap();
3193        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3194
3195        server.abort();
3196    }
3197
3198    #[tokio::test]
3199    async fn malformed_ws_json_reports_parse_error_and_continues() {
3200        let app = Router::new().route("/acp", get(malformed_then_valid_ws));
3201        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3202        let addr = listener.local_addr().unwrap();
3203        let server = tokio::spawn(async move {
3204            axum::serve(listener, app).await.unwrap();
3205        });
3206        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
3207        let (mut caller, transport) = Channel::duplex();
3208        let transport = tokio::spawn(run(client, transport));
3209
3210        let frame = timeout(Duration::from_secs(1), caller.rx.next())
3211            .await
3212            .unwrap()
3213            .unwrap();
3214        let TransportFrame::Malformed { raw, error } = frame else {
3215            panic!("expected malformed frame, got {frame:?}");
3216        };
3217        assert_eq!(raw, "{not json");
3218        assert_eq!(error.code, AcpError::parse_error().code);
3219
3220        let message = timeout(Duration::from_secs(1), caller.rx.next())
3221            .await
3222            .unwrap()
3223            .unwrap()
3224            .unwrap();
3225        assert!(matches!(message, RawJsonRpcMessage::Response(_)));
3226
3227        drop(caller);
3228        timeout(Duration::from_secs(1), transport)
3229            .await
3230            .unwrap()
3231            .unwrap()
3232            .unwrap();
3233
3234        server.abort();
3235    }
3236
3237    #[tokio::test]
3238    async fn websocket_serializes_batch_as_one_text_frame() {
3239        let (caller, transport) = Channel::duplex();
3240        let Channel {
3241            tx: outgoing,
3242            rx: incoming,
3243        } = caller;
3244        drop(incoming);
3245        outgoing
3246            .unbounded_send(TransportFrame::Batch(
3247                TransportBatch::from_messages([
3248                    RawJsonRpcMessage::notification("custom/first".to_string(), json!({})).unwrap(),
3249                    RawJsonRpcMessage::notification("custom/second".to_string(), json!({}))
3250                        .unwrap(),
3251                ])
3252                .unwrap(),
3253            ))
3254            .unwrap();
3255        drop(outgoing);
3256
3257        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
3258        timeout(
3259            Duration::from_secs(1),
3260            drive_ws(
3261                RecordingWsSink(ws_output_tx),
3262                futures::stream::pending::<Result<WsMessage, std::io::Error>>(),
3263                transport,
3264            ),
3265        )
3266        .await
3267        .unwrap()
3268        .unwrap();
3269        let frames = ws_output.by_ref().collect::<Vec<_>>().await;
3270
3271        let WsMessage::Text(text) = &frames[0] else {
3272            panic!("batch was not sent as WebSocket text");
3273        };
3274        let batch = serde_json::from_str::<serde_json::Value>(text.as_str()).unwrap();
3275        let entries = batch.as_array().expect("batch should remain an array");
3276        assert_eq!(entries.len(), 2);
3277        assert_eq!(entries[0]["method"], "custom/first");
3278        assert_eq!(entries[1]["method"], "custom/second");
3279        assert!(matches!(frames.get(1), Some(WsMessage::Close(None))));
3280        assert_eq!(frames.len(), 2);
3281    }
3282
3283    #[tokio::test]
3284    async fn websocket_drain_discards_incoming_after_receiver_closes() {
3285        let (caller, transport) = Channel::duplex();
3286        let Channel {
3287            tx: outgoing,
3288            rx: incoming,
3289        } = caller;
3290        drop(incoming);
3291
3292        let inbound =
3293            RawJsonRpcMessage::notification("custom/inbound".to_string(), json!({})).unwrap();
3294        let inbound = WsMessage::Text(serde_json::to_string(&inbound).unwrap().into());
3295        let ws_rx = QueueOutgoingThenText {
3296            text: Some(inbound),
3297            outgoing: Some(outgoing),
3298        };
3299        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
3300        timeout(
3301            Duration::from_secs(1),
3302            drive_ws(RecordingWsSink(ws_output_tx), ws_rx, transport),
3303        )
3304        .await
3305        .unwrap()
3306        .unwrap();
3307        let mut frames = Vec::new();
3308        while let Some(frame) = ws_output.next().await {
3309            frames.push(frame);
3310        }
3311
3312        let messages = frames
3313            .iter()
3314            .filter_map(|frame| match frame {
3315                WsMessage::Text(text) => {
3316                    Some(serde_json::from_str::<RawJsonRpcMessage>(text.as_str()).unwrap())
3317                }
3318                _ => None,
3319            })
3320            .collect::<Vec<_>>();
3321        let methods = messages
3322            .iter()
3323            .filter_map(method_for_message)
3324            .collect::<Vec<_>>();
3325        assert_eq!(methods, ["custom/first", "custom/second"]);
3326        assert!(matches!(frames.last(), Some(WsMessage::Close(None))));
3327    }
3328
3329    #[tokio::test]
3330    async fn websocket_reader_runs_while_send_is_backpressured() {
3331        let (caller, transport) = Channel::duplex();
3332        let Channel {
3333            tx: outgoing,
3334            rx: incoming,
3335        } = caller;
3336        drop(incoming);
3337        outgoing
3338            .unbounded_send(single_frame(
3339                RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(),
3340            ))
3341            .unwrap();
3342        drop(outgoing);
3343
3344        let (started_tx, started_rx) = mpsc::unbounded();
3345        let (release_tx, release_rx) = futures::channel::oneshot::channel();
3346        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
3347        let ws_tx = BackpressuredWsSink {
3348            output: ws_output_tx,
3349            started: started_tx,
3350            release: Some(release_rx),
3351        };
3352        let ws_rx = ReleaseBackpressureOnPoll {
3353            started: started_rx,
3354            release: Some(release_tx),
3355        };
3356
3357        timeout(Duration::from_secs(1), drive_ws(ws_tx, ws_rx, transport))
3358            .await
3359            .expect("WebSocket reader was not polled while its writer was backpressured")
3360            .unwrap();
3361        let frames = ws_output.by_ref().collect::<Vec<_>>().await;
3362
3363        let WsMessage::Text(text) = &frames[0] else {
3364            panic!("queued message was not sent as WebSocket text");
3365        };
3366        let message = serde_json::from_str::<RawJsonRpcMessage>(text.as_str()).unwrap();
3367        assert_eq!(method_for_message(&message), Some("custom/queued"));
3368        assert!(matches!(frames.get(1), Some(WsMessage::Close(None))));
3369        assert_eq!(frames.len(), 2);
3370    }
3371
3372    #[tokio::test]
3373    async fn peer_ws_close_fails_transport() {
3374        let app = Router::new().route("/acp", get(close_ws));
3375        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3376        let addr = listener.local_addr().unwrap();
3377        let server = tokio::spawn(async move {
3378            axum::serve(listener, app).await.unwrap();
3379        });
3380        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
3381        let (_caller, transport) = Channel::duplex();
3382        let transport = tokio::spawn(run(client, transport));
3383
3384        let error = timeout(Duration::from_secs(1), transport)
3385            .await
3386            .unwrap()
3387            .unwrap()
3388            .unwrap_err();
3389        assert!(error.to_string().contains("WebSocket closed by peer"));
3390
3391        server.abort();
3392    }
3393
3394    #[tokio::test]
3395    async fn dropped_transport_future_deletes_initialized_connection() {
3396        let delete_count = Arc::new(AtomicUsize::new(0));
3397        let delete_count_for_handler = delete_count.clone();
3398        let app = Router::new().route(
3399            "/acp",
3400            post(initialize_response).get(pending_sse).delete(move || {
3401                let delete_count = delete_count_for_handler.clone();
3402                async move {
3403                    delete_count.fetch_add(1, Ordering::SeqCst);
3404                    StatusCode::ACCEPTED
3405                }
3406            }),
3407        );
3408        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3409        let addr = listener.local_addr().unwrap();
3410        let server = tokio::spawn(async move {
3411            axum::serve(listener, app).await.unwrap();
3412        });
3413        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3414        let (mut caller, transport) = Channel::duplex();
3415        let mut transport = Box::pin(run(client, transport));
3416
3417        caller
3418            .tx
3419            .unbounded_send(single_frame(
3420                RawJsonRpcMessage::request(
3421                    "initialize".to_string(),
3422                    json!({}),
3423                    RequestId::Number(1),
3424                )
3425                .unwrap(),
3426            ))
3427            .unwrap();
3428        let init_response = timeout(Duration::from_secs(1), async {
3429            tokio::select! {
3430                result = &mut transport => {
3431                    panic!("transport ended before initialize response: {result:?}");
3432                }
3433                msg = caller.rx.next() => {
3434                    msg.unwrap().unwrap()
3435                }
3436            }
3437        })
3438        .await
3439        .unwrap();
3440        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3441
3442        drop(transport);
3443        wait_for_delete(&delete_count).await;
3444
3445        server.abort();
3446    }
3447
3448    #[tokio::test]
3449    async fn dropped_transport_during_close_retries_delete() {
3450        let delete_count = Arc::new(AtomicUsize::new(0));
3451        let delete_count_for_handler = delete_count.clone();
3452        let release_delete = Arc::new(Notify::new());
3453        let release_delete_for_handler = release_delete.clone();
3454        let app = Router::new().route(
3455            "/acp",
3456            post(initialize_response).get(pending_sse).delete(move || {
3457                let delete_count = delete_count_for_handler.clone();
3458                let release_delete = release_delete_for_handler.clone();
3459                async move {
3460                    delete_count.fetch_add(1, Ordering::SeqCst);
3461                    release_delete.notified().await;
3462                    StatusCode::ACCEPTED
3463                }
3464            }),
3465        );
3466        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3467        let addr = listener.local_addr().unwrap();
3468        let server = tokio::spawn(async move {
3469            axum::serve(listener, app).await.unwrap();
3470        });
3471        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3472        let (mut caller, transport) = Channel::duplex();
3473        let transport = tokio::spawn(run(client, transport));
3474
3475        caller
3476            .tx
3477            .unbounded_send(single_frame(
3478                RawJsonRpcMessage::request(
3479                    "initialize".to_string(),
3480                    json!({}),
3481                    RequestId::Number(1),
3482                )
3483                .unwrap(),
3484            ))
3485            .unwrap();
3486        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3487            .await
3488            .unwrap()
3489            .unwrap()
3490            .unwrap();
3491        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3492
3493        drop(caller);
3494        wait_for_delete_count(&delete_count, 1).await;
3495        transport.abort();
3496        wait_for_delete_count(&delete_count, 2).await;
3497        release_delete.notify_waiters();
3498        drop(transport.await);
3499
3500        server.abort();
3501    }
3502
3503    #[tokio::test]
3504    async fn initialize_error_without_connection_id_is_delivered_without_sse() {
3505        let get_count = Arc::new(AtomicUsize::new(0));
3506        let get_count_for_handler = get_count.clone();
3507        let app = Router::new().route(
3508            "/acp",
3509            post(initialize_error_response).get(move || {
3510                let get_count = get_count_for_handler.clone();
3511                async move {
3512                    get_count.fetch_add(1, Ordering::SeqCst);
3513                    pending_sse().await
3514                }
3515            }),
3516        );
3517        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3518        let addr = listener.local_addr().unwrap();
3519        let server = tokio::spawn(async move {
3520            axum::serve(listener, app).await.unwrap();
3521        });
3522        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3523        let (mut caller, transport) = Channel::duplex();
3524        let transport = tokio::spawn(run(client, transport));
3525
3526        caller
3527            .tx
3528            .unbounded_send(single_frame(
3529                RawJsonRpcMessage::request(
3530                    "initialize".to_string(),
3531                    json!({}),
3532                    RequestId::Number(1),
3533                )
3534                .unwrap(),
3535            ))
3536            .unwrap();
3537        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3538            .await
3539            .unwrap()
3540            .unwrap()
3541            .unwrap();
3542
3543        assert!(matches!(
3544            init_response,
3545            RawJsonRpcMessage::Response(RpcResponse::Error {
3546                id: RequestId::Number(1),
3547                ..
3548            })
3549        ));
3550        assert_eq!(get_count.load(Ordering::SeqCst), 0);
3551
3552        drop(caller);
3553        timeout(Duration::from_secs(1), transport)
3554            .await
3555            .unwrap()
3556            .unwrap()
3557            .unwrap();
3558
3559        server.abort();
3560    }
3561
3562    #[tokio::test]
3563    async fn malformed_initialize_body_with_connection_id_is_deleted() {
3564        let delete_count = Arc::new(AtomicUsize::new(0));
3565        let delete_count_for_handler = delete_count.clone();
3566        let app = Router::new().route(
3567            "/acp",
3568            post(malformed_initialize_response).delete(move || {
3569                let delete_count = delete_count_for_handler.clone();
3570                async move {
3571                    delete_count.fetch_add(1, Ordering::SeqCst);
3572                    StatusCode::ACCEPTED
3573                }
3574            }),
3575        );
3576        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3577        let addr = listener.local_addr().unwrap();
3578        let server = tokio::spawn(async move {
3579            axum::serve(listener, app).await.unwrap();
3580        });
3581        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3582        let (caller, transport) = Channel::duplex();
3583        let transport = tokio::spawn(run(client, transport));
3584
3585        caller
3586            .tx
3587            .unbounded_send(single_frame(
3588                RawJsonRpcMessage::request(
3589                    "initialize".to_string(),
3590                    json!({}),
3591                    RequestId::Number(1),
3592                )
3593                .unwrap(),
3594            ))
3595            .unwrap();
3596        let error = timeout(Duration::from_secs(1), transport)
3597            .await
3598            .unwrap()
3599            .unwrap()
3600            .unwrap_err();
3601
3602        assert!(error.to_string().contains("initialize"));
3603        wait_for_delete(&delete_count).await;
3604
3605        server.abort();
3606    }
3607
3608    async fn wait_for_delete(delete_count: &AtomicUsize) {
3609        wait_for_delete_count(delete_count, 1).await;
3610        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3611    }
3612
3613    async fn wait_for_delete_count(delete_count: &AtomicUsize, expected: usize) {
3614        timeout(Duration::from_secs(1), async {
3615            loop {
3616                if delete_count.load(Ordering::SeqCst) >= expected {
3617                    break;
3618                }
3619                sleep(Duration::from_millis(10)).await;
3620            }
3621        })
3622        .await
3623        .unwrap();
3624    }
3625
3626    async fn initialize_response() -> impl IntoResponse {
3627        let mut headers = HeaderMap::new();
3628        headers.insert(HEADER_CONNECTION_ID, HeaderValue::from_static("conn-1"));
3629        (
3630            StatusCode::OK,
3631            headers,
3632            Json(RawJsonRpcMessage::response(
3633                RequestId::Number(1),
3634                Ok(json!({})),
3635            )),
3636        )
3637    }
3638
3639    async fn initialize_error_response() -> Json<RawJsonRpcMessage> {
3640        Json(RawJsonRpcMessage::response(
3641            RequestId::Number(1),
3642            Err(AcpError::invalid_request().data("initialize rejected")),
3643        ))
3644    }
3645
3646    async fn malformed_initialize_response() -> impl IntoResponse {
3647        let mut headers = HeaderMap::new();
3648        headers.insert(HEADER_CONNECTION_ID, HeaderValue::from_static("conn-1"));
3649        (StatusCode::OK, headers, "{not json")
3650    }
3651
3652    async fn pending_sse() -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
3653        Sse::new(futures::stream::pending())
3654    }
3655
3656    fn sse_event(message: RawJsonRpcMessage) -> Event {
3657        Event::default().data(serde_json::to_string(&message).unwrap())
3658    }
3659
3660    async fn malformed_sse() -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
3661        let invalid = futures::stream::once(async {
3662            Ok::<_, Infallible>(Event::default().data("{not json"))
3663        });
3664        Sse::new(invalid.chain(futures::stream::pending()))
3665    }
3666
3667    async fn malformed_then_valid_ws(ws: WebSocketUpgrade) -> impl IntoResponse {
3668        ws.on_upgrade(|mut socket| async move {
3669            drop(socket.send(AxumWsMessage::Text("{not json".into())).await);
3670            let valid = serde_json::to_string(&RawJsonRpcMessage::response(
3671                RequestId::Number(1),
3672                Ok(json!({})),
3673            ))
3674            .unwrap();
3675            drop(socket.send(AxumWsMessage::Text(valid.into())).await);
3676            futures::future::pending::<()>().await;
3677        })
3678    }
3679
3680    async fn close_ws(ws: WebSocketUpgrade) -> impl IntoResponse {
3681        ws.on_upgrade(|mut socket| async move {
3682            drop(socket.send(AxumWsMessage::Close(None)).await);
3683        })
3684    }
3685
3686    async fn closed_sse() -> StatusCode {
3687        StatusCode::OK
3688    }
3689}