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        async fn send(&mut self, message: WsMessage) -> Result<(), String> {
1477            self.0
1478                .unbounded_send(message)
1479                .map_err(|error| error.to_string())
1480        }
1481    }
1482
1483    impl WsSink for BackpressuredWsSink {
1484        async fn send(&mut self, message: WsMessage) -> Result<(), String> {
1485            self.output
1486                .unbounded_send(message)
1487                .map_err(|error| error.to_string())?;
1488            if let Some(release) = self.release.take() {
1489                self.started
1490                    .unbounded_send(())
1491                    .map_err(|error| error.to_string())?;
1492                release
1493                    .await
1494                    .map_err(|_| "mock WebSocket reader did not release send".to_string())?;
1495            }
1496            Ok(())
1497        }
1498    }
1499
1500    impl Stream for QueueOutgoingThenText {
1501        type Item = Result<WsMessage, std::io::Error>;
1502
1503        fn poll_next(
1504            mut self: std::pin::Pin<&mut Self>,
1505            _cx: &mut std::task::Context<'_>,
1506        ) -> std::task::Poll<Option<Self::Item>> {
1507            // Make input ready immediately after queueing output. If the output
1508            // branch was polled first it was still empty, so either poll order
1509            // deterministically selects this input frame first.
1510            if let Some(outgoing) = self.outgoing.take() {
1511                for method in ["custom/first", "custom/second"] {
1512                    outgoing
1513                        .unbounded_send(single_frame(
1514                            RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(),
1515                        ))
1516                        .unwrap();
1517                }
1518            }
1519            if let Some(text) = self.text.take() {
1520                return std::task::Poll::Ready(Some(Ok(text)));
1521            }
1522            std::task::Poll::Pending
1523        }
1524    }
1525
1526    impl Stream for ReleaseBackpressureOnPoll {
1527        type Item = Result<WsMessage, std::io::Error>;
1528
1529        fn poll_next(
1530            mut self: std::pin::Pin<&mut Self>,
1531            cx: &mut std::task::Context<'_>,
1532        ) -> std::task::Poll<Option<Self::Item>> {
1533            if let std::task::Poll::Ready(Some(())) =
1534                std::pin::Pin::new(&mut self.started).poll_next(cx)
1535                && let Some(release) = self.release.take()
1536            {
1537                let _result = release.send(());
1538            }
1539            std::task::Poll::Pending
1540        }
1541    }
1542
1543    impl ConnectTo<Agent> for PostsThenExitClient {
1544        async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), AcpError> {
1545            let Self {
1546                finish,
1547                finished,
1548                escaped_tx,
1549            } = self;
1550            let (mut channel, transport) = agent.into_channel_and_future();
1551            let client = async move {
1552                escaped_tx.send(channel.tx.clone()).map_err(|_| {
1553                    AcpError::internal_error().data("escaped sender observer dropped")
1554                })?;
1555                channel
1556                    .tx
1557                    .unbounded_send(single_frame(
1558                        RawJsonRpcMessage::request(
1559                            "initialize".to_string(),
1560                            json!({}),
1561                            RequestId::Number(1),
1562                        )
1563                        .unwrap(),
1564                    ))
1565                    .map_err(|e| {
1566                        AcpError::internal_error().data(format!("send initialize: {e}"))
1567                    })?;
1568                into_single_message(channel.rx.next().await.ok_or_else(|| {
1569                    AcpError::internal_error().data("initialize response channel closed")
1570                })?)?;
1571
1572                for method in ["custom/first", "custom/second"] {
1573                    channel
1574                        .tx
1575                        .unbounded_send(single_frame(
1576                            RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(),
1577                        ))
1578                        .map_err(|e| {
1579                            AcpError::internal_error().data(format!("send {method}: {e}"))
1580                        })?;
1581                }
1582
1583                finish.notified().await;
1584                finished.notify_one();
1585                Ok(())
1586            };
1587
1588            let ((), ()) = futures::try_join!(transport, client)?;
1589            Ok(())
1590        }
1591    }
1592
1593    impl ConnectTo<Agent> for InitializeThenExitClient {
1594        async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), AcpError> {
1595            let Self {
1596                sse_started,
1597                finished,
1598            } = self;
1599            let (mut channel, transport) = agent.into_channel_and_future();
1600            let client = async move {
1601                channel
1602                    .tx
1603                    .unbounded_send(single_frame(
1604                        RawJsonRpcMessage::request(
1605                            "initialize".to_string(),
1606                            json!({}),
1607                            RequestId::Number(1),
1608                        )
1609                        .unwrap(),
1610                    ))
1611                    .map_err(|error| {
1612                        AcpError::internal_error().data(format!("send initialize: {error}"))
1613                    })?;
1614                into_single_message(channel.rx.next().await.ok_or_else(|| {
1615                    AcpError::internal_error().data("initialize response channel closed")
1616                })?)?;
1617
1618                sse_started.notified().await;
1619                finished.notify_one();
1620                Ok(())
1621            };
1622
1623            let ((), ()) = futures::try_join!(transport, client)?;
1624            Ok(())
1625        }
1626    }
1627
1628    #[test]
1629    fn new_targets_standard_acp_endpoint() {
1630        assert_eq!(
1631            HttpClient::new("http://example.com")
1632                .unwrap()
1633                .endpoint
1634                .as_str(),
1635            "http://example.com/acp"
1636        );
1637        assert_eq!(
1638            HttpClient::new("http://example.com/proxy")
1639                .unwrap()
1640                .endpoint
1641                .as_str(),
1642            "http://example.com/proxy/acp"
1643        );
1644        assert_eq!(
1645            HttpClient::new("http://example.com/proxy/acp")
1646                .unwrap()
1647                .endpoint
1648                .as_str(),
1649            "http://example.com/proxy/acp"
1650        );
1651    }
1652
1653    #[test]
1654    fn with_endpoint_preserves_explicit_endpoint_path() {
1655        assert_eq!(
1656            HttpClient::with_endpoint("http://example.com/agent")
1657                .unwrap()
1658                .endpoint
1659                .as_str(),
1660            "http://example.com/agent"
1661        );
1662        assert_eq!(
1663            HttpClient::with_endpoint_and_client(
1664                "ws://example.com/custom/acp?token=abc",
1665                reqwest::Client::new(),
1666            )
1667            .unwrap()
1668            .endpoint
1669            .as_str(),
1670            "ws://example.com/custom/acp?token=abc"
1671        );
1672    }
1673
1674    #[tokio::test]
1675    async fn post_sends_cancel_request_without_session_header() {
1676        let (capture_tx, mut capture_rx) = tokio::sync::mpsc::unbounded_channel();
1677        let post_count = Arc::new(AtomicUsize::new(0));
1678        let app = Router::new().route(
1679            "/acp",
1680            post({
1681                let capture_tx = capture_tx.clone();
1682                let post_count = post_count.clone();
1683                move |headers: HeaderMap, Json(message): Json<RawJsonRpcMessage>| {
1684                    let capture_tx = capture_tx.clone();
1685                    let post_count = post_count.clone();
1686                    async move {
1687                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1688                            return initialize_response().await.into_response();
1689                        }
1690
1691                        capture_tx
1692                            .send((headers.get(HEADER_SESSION_ID).cloned(), message))
1693                            .unwrap();
1694                        StatusCode::ACCEPTED.into_response()
1695                    }
1696                }
1697            })
1698            .get(pending_sse)
1699            .delete(|| async { StatusCode::ACCEPTED }),
1700        );
1701        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1702        let addr = listener.local_addr().unwrap();
1703        let server = tokio::spawn(async move {
1704            axum::serve(listener, app).await.unwrap();
1705        });
1706        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1707        let (mut caller, transport) = Channel::duplex();
1708        let transport = tokio::spawn(run(client, transport));
1709
1710        caller
1711            .tx
1712            .unbounded_send(single_frame(
1713                RawJsonRpcMessage::request(
1714                    "initialize".to_string(),
1715                    json!({}),
1716                    RequestId::Number(1),
1717                )
1718                .unwrap(),
1719            ))
1720            .unwrap();
1721        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1722            .await
1723            .unwrap()
1724            .unwrap()
1725            .unwrap();
1726        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1727
1728        caller
1729            .tx
1730            .unbounded_send(single_frame(
1731                RawJsonRpcMessage::notification(
1732                    "$/cancel_request".to_string(),
1733                    json!({
1734                        "requestId": 2,
1735                        "sessionId": "session-1"
1736                    }),
1737                )
1738                .unwrap(),
1739            ))
1740            .unwrap();
1741
1742        let (session_header, message) = timeout(Duration::from_secs(1), capture_rx.recv())
1743            .await
1744            .unwrap()
1745            .unwrap();
1746        assert!(session_header.is_none());
1747        assert!(matches!(
1748            message,
1749            RawJsonRpcMessage::Notification(notification)
1750                if notification.method.as_ref() == "$/cancel_request"
1751        ));
1752
1753        drop(caller);
1754        timeout(Duration::from_secs(1), transport)
1755            .await
1756            .unwrap()
1757            .unwrap()
1758            .unwrap();
1759
1760        server.abort();
1761    }
1762
1763    #[tokio::test]
1764    async fn http_preserves_batch_frames_across_post_and_sse() {
1765        let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel();
1766        let post_count = Arc::new(AtomicUsize::new(0));
1767        let emit_sse = Arc::new(Notify::new());
1768        let inbound_batch = json!([
1769            {
1770                "jsonrpc": "2.0",
1771                "method": "custom/inbound-one",
1772                "params": {}
1773            },
1774            {
1775                "jsonrpc": "2.0",
1776                "method": "custom/inbound-two",
1777                "params": {}
1778            }
1779        ]);
1780        let app = Router::new().route(
1781            "/acp",
1782            post({
1783                let post_count = post_count.clone();
1784                move |body: String| {
1785                    let post_count = post_count.clone();
1786                    let post_tx = post_tx.clone();
1787                    async move {
1788                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1789                            return initialize_response().await.into_response();
1790                        }
1791
1792                        post_tx
1793                            .send(serde_json::from_str::<serde_json::Value>(&body).unwrap())
1794                            .unwrap();
1795                        StatusCode::ACCEPTED.into_response()
1796                    }
1797                }
1798            })
1799            .get({
1800                let emit_sse = emit_sse.clone();
1801                let inbound_batch = inbound_batch.clone();
1802                move || {
1803                    let emit_sse = emit_sse.clone();
1804                    let inbound_batch = inbound_batch.clone();
1805                    async move {
1806                        let stream = async_stream::stream! {
1807                            emit_sse.notified().await;
1808                            yield Ok::<_, Infallible>(
1809                                Event::default().data(inbound_batch.to_string()),
1810                            );
1811                            futures::future::pending::<()>().await;
1812                        };
1813                        Sse::new(stream)
1814                    }
1815                }
1816            })
1817            .delete(|| async { StatusCode::ACCEPTED }),
1818        );
1819        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1820        let addr = listener.local_addr().unwrap();
1821        let server = tokio::spawn(async move {
1822            axum::serve(listener, app).await.unwrap();
1823        });
1824        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1825        let (mut caller, transport) = Channel::duplex();
1826        let transport = tokio::spawn(run(client, transport));
1827
1828        caller
1829            .tx
1830            .unbounded_send(single_frame(
1831                RawJsonRpcMessage::request(
1832                    "initialize".to_string(),
1833                    json!({}),
1834                    RequestId::Number(1),
1835                )
1836                .unwrap(),
1837            ))
1838            .unwrap();
1839        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1840            .await
1841            .unwrap()
1842            .unwrap()
1843            .unwrap();
1844        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1845
1846        let outbound_batch = json!([
1847            {
1848                "jsonrpc": "2.0",
1849                "method": "custom/outbound-one",
1850                "params": {}
1851            },
1852            {
1853                "jsonrpc": "2.0",
1854                "method": "custom/outbound-two",
1855                "params": {}
1856            }
1857        ]);
1858        caller
1859            .tx
1860            .unbounded_send(TransportFrame::Batch(
1861                TransportBatch::from_messages([
1862                    RawJsonRpcMessage::notification("custom/outbound-one".to_string(), json!({}))
1863                        .unwrap(),
1864                    RawJsonRpcMessage::notification("custom/outbound-two".to_string(), json!({}))
1865                        .unwrap(),
1866                ])
1867                .unwrap(),
1868            ))
1869            .unwrap();
1870
1871        let posted = timeout(Duration::from_secs(1), post_rx.recv())
1872            .await
1873            .unwrap()
1874            .unwrap();
1875        assert_eq!(posted, outbound_batch);
1876
1877        emit_sse.notify_one();
1878        let inbound = timeout(Duration::from_secs(1), caller.rx.next())
1879            .await
1880            .unwrap()
1881            .unwrap();
1882        assert!(matches!(&inbound, TransportFrame::Batch(_)));
1883        assert_eq!(
1884            serde_json::from_str::<serde_json::Value>(&inbound.to_json().unwrap()).unwrap(),
1885            inbound_batch
1886        );
1887
1888        drop(caller);
1889        timeout(Duration::from_secs(1), transport)
1890            .await
1891            .unwrap()
1892            .unwrap()
1893            .unwrap();
1894
1895        server.abort();
1896    }
1897
1898    #[tokio::test]
1899    async fn batch_fork_opens_source_and_result_session_streams() {
1900        let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel();
1901        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
1902        let post_count = Arc::new(AtomicUsize::new(0));
1903        let emit_response = Arc::new(Notify::new());
1904        let connection_stream_established = Arc::new(AtomicBool::new(false));
1905        let source_stream_established = Arc::new(AtomicBool::new(false));
1906        let response_batch = json!([
1907            {
1908                "jsonrpc": "2.0",
1909                "id": 2,
1910                "result": { "sessionId": "forked-session" }
1911            }
1912        ]);
1913        let app = Router::new().route(
1914            "/acp",
1915            post({
1916                let post_count = post_count.clone();
1917                let connection_stream_established = connection_stream_established.clone();
1918                let source_stream_established = source_stream_established.clone();
1919                move |body: String| {
1920                    let post_count = post_count.clone();
1921                    let post_tx = post_tx.clone();
1922                    let connection_stream_established = connection_stream_established.clone();
1923                    let source_stream_established = source_stream_established.clone();
1924                    async move {
1925                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1926                            return initialize_response().await.into_response();
1927                        }
1928
1929                        if !connection_stream_established.load(Ordering::SeqCst)
1930                            || !source_stream_established.load(Ordering::SeqCst)
1931                        {
1932                            return StatusCode::CONFLICT.into_response();
1933                        }
1934                        post_tx
1935                            .send(serde_json::from_str::<serde_json::Value>(&body).unwrap())
1936                            .unwrap();
1937                        StatusCode::ACCEPTED.into_response()
1938                    }
1939                }
1940            })
1941            .get({
1942                let emit_response = emit_response.clone();
1943                let response_batch = response_batch.clone();
1944                let connection_stream_established = connection_stream_established.clone();
1945                let source_stream_established = source_stream_established.clone();
1946                move |headers: HeaderMap| {
1947                    let emit_response = emit_response.clone();
1948                    let response_batch = response_batch.clone();
1949                    let get_tx = get_tx.clone();
1950                    let connection_stream_established = connection_stream_established.clone();
1951                    let source_stream_established = source_stream_established.clone();
1952                    async move {
1953                        let session_id = headers
1954                            .get(HEADER_SESSION_ID)
1955                            .and_then(|value| value.to_str().ok())
1956                            .map(String::from);
1957                        let is_connection_stream = session_id.is_none();
1958                        let is_source_stream = session_id.as_deref() == Some("source-session");
1959                        if is_connection_stream {
1960                            sleep(Duration::from_millis(50)).await;
1961                            connection_stream_established.store(true, Ordering::SeqCst);
1962                        }
1963                        if is_source_stream {
1964                            sleep(Duration::from_millis(50)).await;
1965                            source_stream_established.store(true, Ordering::SeqCst);
1966                        }
1967                        get_tx.send(session_id).unwrap();
1968
1969                        let stream = async_stream::stream! {
1970                            if is_source_stream {
1971                                emit_response.notified().await;
1972                                yield Ok::<_, Infallible>(
1973                                    Event::default().data(response_batch.to_string()),
1974                                );
1975                            }
1976                            futures::future::pending::<()>().await;
1977                        };
1978                        Sse::new(stream)
1979                    }
1980                }
1981            })
1982            .delete(|| async { StatusCode::ACCEPTED }),
1983        );
1984        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1985        let addr = listener.local_addr().unwrap();
1986        let server = tokio::spawn(async move {
1987            axum::serve(listener, app).await.unwrap();
1988        });
1989        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1990        let (mut caller, transport) = Channel::duplex();
1991        let transport = tokio::spawn(run(client, transport));
1992
1993        caller
1994            .tx
1995            .unbounded_send(single_frame(
1996                RawJsonRpcMessage::request(
1997                    "initialize".to_string(),
1998                    json!({}),
1999                    RequestId::Number(1),
2000                )
2001                .unwrap(),
2002            ))
2003            .unwrap();
2004        timeout(Duration::from_secs(1), caller.rx.next())
2005            .await
2006            .unwrap()
2007            .unwrap();
2008
2009        caller
2010            .tx
2011            .unbounded_send(TransportFrame::Batch(
2012                TransportBatch::from_messages([RawJsonRpcMessage::request(
2013                    "session/fork".to_string(),
2014                    json!({ "sessionId": "source-session" }),
2015                    RequestId::Number(2),
2016                )
2017                .unwrap()])
2018                .unwrap(),
2019            ))
2020            .unwrap();
2021
2022        let connection_stream = timeout(Duration::from_secs(1), get_rx.recv())
2023            .await
2024            .unwrap()
2025            .unwrap();
2026        assert!(connection_stream.is_none());
2027        let source_stream = timeout(Duration::from_secs(1), get_rx.recv())
2028            .await
2029            .unwrap()
2030            .unwrap();
2031        assert_eq!(source_stream.as_deref(), Some("source-session"));
2032        let posted = timeout(Duration::from_secs(1), post_rx.recv())
2033            .await
2034            .unwrap()
2035            .unwrap();
2036        assert!(posted.is_array(), "outgoing batch must remain an array");
2037
2038        emit_response.notify_one();
2039        let response = timeout(Duration::from_secs(1), caller.rx.next())
2040            .await
2041            .unwrap()
2042            .unwrap();
2043        assert!(matches!(&response, TransportFrame::Batch(_)));
2044        assert_eq!(
2045            serde_json::from_str::<serde_json::Value>(&response.to_json().unwrap()).unwrap(),
2046            response_batch
2047        );
2048        let forked_stream = timeout(Duration::from_secs(1), get_rx.recv())
2049            .await
2050            .unwrap()
2051            .unwrap();
2052        assert_eq!(forked_stream.as_deref(), Some("forked-session"));
2053
2054        drop(caller);
2055        timeout(Duration::from_secs(1), transport)
2056            .await
2057            .unwrap()
2058            .unwrap()
2059            .unwrap();
2060
2061        server.abort();
2062    }
2063
2064    #[tokio::test]
2065    async fn custom_response_with_session_id_does_not_open_session_sse() {
2066        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
2067        let response_ready = Arc::new(tokio::sync::Notify::new());
2068        let post_count = Arc::new(AtomicUsize::new(0));
2069        let app = Router::new().route(
2070            "/acp",
2071            post({
2072                let post_count = post_count.clone();
2073                let response_ready = response_ready.clone();
2074                move |Json(_message): Json<RawJsonRpcMessage>| {
2075                    let post_count = post_count.clone();
2076                    let response_ready = response_ready.clone();
2077                    async move {
2078                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2079                            return initialize_response().await.into_response();
2080                        }
2081
2082                        response_ready.notify_waiters();
2083                        StatusCode::ACCEPTED.into_response()
2084                    }
2085                }
2086            })
2087            .get({
2088                let get_tx = get_tx.clone();
2089                let response_ready = response_ready.clone();
2090                move |headers: HeaderMap| {
2091                    let get_tx = get_tx.clone();
2092                    let response_ready = response_ready.clone();
2093                    async move {
2094                        let session_header = headers
2095                            .get(HEADER_SESSION_ID)
2096                            .and_then(|value| value.to_str().ok())
2097                            .map(String::from);
2098                        get_tx.send(session_header).unwrap();
2099
2100                        let stream = async_stream::stream! {
2101                            response_ready.notified().await;
2102                            yield Ok::<_, Infallible>(sse_event(
2103                                RawJsonRpcMessage::response(
2104                                    RequestId::Number(2),
2105                                    Ok(json!({ "sessionId": "session-1" })),
2106                                ),
2107                            ));
2108                            futures::future::pending::<()>().await;
2109                        };
2110                        Sse::new(stream)
2111                    }
2112                }
2113            })
2114            .delete(|| async { StatusCode::ACCEPTED }),
2115        );
2116        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2117        let addr = listener.local_addr().unwrap();
2118        let server = tokio::spawn(async move {
2119            axum::serve(listener, app).await.unwrap();
2120        });
2121        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2122        let (mut caller, transport) = Channel::duplex();
2123        let transport = tokio::spawn(run(client, transport));
2124
2125        caller
2126            .tx
2127            .unbounded_send(single_frame(
2128                RawJsonRpcMessage::request(
2129                    "initialize".to_string(),
2130                    json!({}),
2131                    RequestId::Number(1),
2132                )
2133                .unwrap(),
2134            ))
2135            .unwrap();
2136        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2137            .await
2138            .unwrap()
2139            .unwrap()
2140            .unwrap();
2141        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2142
2143        let connection_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2144            .await
2145            .unwrap()
2146            .unwrap();
2147        assert!(connection_sse_header.is_none());
2148
2149        caller
2150            .tx
2151            .unbounded_send(single_frame(
2152                RawJsonRpcMessage::request(
2153                    "custom/sessionish".to_string(),
2154                    json!({}),
2155                    RequestId::Number(2),
2156                )
2157                .unwrap(),
2158            ))
2159            .unwrap();
2160        let response = timeout(Duration::from_secs(1), caller.rx.next())
2161            .await
2162            .unwrap()
2163            .unwrap()
2164            .unwrap();
2165        assert!(matches!(
2166            response,
2167            RawJsonRpcMessage::Response(RpcResponse::Result {
2168                id: RequestId::Number(2),
2169                ..
2170            })
2171        ));
2172
2173        assert!(
2174            timeout(Duration::from_millis(100), get_rx.recv())
2175                .await
2176                .is_err(),
2177            "custom response must not open a session SSE stream"
2178        );
2179
2180        drop(caller);
2181        timeout(Duration::from_secs(1), transport)
2182            .await
2183            .unwrap()
2184            .unwrap()
2185            .unwrap();
2186
2187        server.abort();
2188    }
2189
2190    #[tokio::test]
2191    async fn fork_response_with_session_id_opens_session_sse() {
2192        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
2193        let response_ready = Arc::new(tokio::sync::Notify::new());
2194        let post_count = Arc::new(AtomicUsize::new(0));
2195        let app = Router::new().route(
2196            "/acp",
2197            post({
2198                let post_count = post_count.clone();
2199                let response_ready = response_ready.clone();
2200                move |Json(_message): Json<RawJsonRpcMessage>| {
2201                    let post_count = post_count.clone();
2202                    let response_ready = response_ready.clone();
2203                    async move {
2204                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2205                            return initialize_response().await.into_response();
2206                        }
2207
2208                        response_ready.notify_waiters();
2209                        StatusCode::ACCEPTED.into_response()
2210                    }
2211                }
2212            })
2213            .get({
2214                let get_tx = get_tx.clone();
2215                let response_ready = response_ready.clone();
2216                move |headers: HeaderMap| {
2217                    let get_tx = get_tx.clone();
2218                    let response_ready = response_ready.clone();
2219                    async move {
2220                        let session_header = headers
2221                            .get(HEADER_SESSION_ID)
2222                            .and_then(|value| value.to_str().ok())
2223                            .map(String::from);
2224                        let is_connection_stream = session_header.is_none();
2225                        get_tx.send(session_header).unwrap();
2226
2227                        let stream = async_stream::stream! {
2228                            if is_connection_stream {
2229                                response_ready.notified().await;
2230                                yield Ok::<_, Infallible>(sse_event(
2231                                    RawJsonRpcMessage::response(
2232                                        RequestId::Number(2),
2233                                        Ok(json!({ "sessionId": "forked-session" })),
2234                                    ),
2235                                ));
2236                            }
2237                            futures::future::pending::<()>().await;
2238                        };
2239                        Sse::new(stream)
2240                    }
2241                }
2242            })
2243            .delete(|| async { StatusCode::ACCEPTED }),
2244        );
2245        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2246        let addr = listener.local_addr().unwrap();
2247        let server = tokio::spawn(async move {
2248            axum::serve(listener, app).await.unwrap();
2249        });
2250        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2251        let (mut caller, transport) = Channel::duplex();
2252        let transport = tokio::spawn(run(client, transport));
2253
2254        caller
2255            .tx
2256            .unbounded_send(single_frame(
2257                RawJsonRpcMessage::request(
2258                    "initialize".to_string(),
2259                    json!({}),
2260                    RequestId::Number(1),
2261                )
2262                .unwrap(),
2263            ))
2264            .unwrap();
2265        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2266            .await
2267            .unwrap()
2268            .unwrap()
2269            .unwrap();
2270        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2271
2272        let connection_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2273            .await
2274            .unwrap()
2275            .unwrap();
2276        assert!(connection_sse_header.is_none());
2277
2278        caller
2279            .tx
2280            .unbounded_send(single_frame(
2281                RawJsonRpcMessage::request(
2282                    "session/fork".to_string(),
2283                    json!({ "sessionId": "source-session" }),
2284                    RequestId::Number(2),
2285                )
2286                .unwrap(),
2287            ))
2288            .unwrap();
2289        let source_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2290            .await
2291            .unwrap()
2292            .unwrap();
2293        assert_eq!(source_sse_header.as_deref(), Some("source-session"));
2294
2295        let response = timeout(Duration::from_secs(1), caller.rx.next())
2296            .await
2297            .unwrap()
2298            .unwrap()
2299            .unwrap();
2300        assert!(matches!(
2301            response,
2302            RawJsonRpcMessage::Response(RpcResponse::Result {
2303                id: RequestId::Number(2),
2304                ..
2305            })
2306        ));
2307
2308        let fork_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2309            .await
2310            .unwrap()
2311            .unwrap();
2312        assert_eq!(fork_sse_header.as_deref(), Some("forked-session"));
2313
2314        drop(caller);
2315        timeout(Duration::from_secs(1), transport)
2316            .await
2317            .unwrap()
2318            .unwrap()
2319            .unwrap();
2320
2321        server.abort();
2322    }
2323
2324    #[tokio::test]
2325    async fn only_response_batches_bypass_ordered_posts() {
2326        let slow_started = Arc::new(Notify::new());
2327        let release_slow = Arc::new(Notify::new());
2328        let call_batch_seen = Arc::new(Notify::new());
2329        let response_batch_seen = Arc::new(Notify::new());
2330        let app = Router::new().route(
2331            "/acp",
2332            post({
2333                let slow_started = slow_started.clone();
2334                let release_slow = release_slow.clone();
2335                let call_batch_seen = call_batch_seen.clone();
2336                let response_batch_seen = response_batch_seen.clone();
2337                move |body: String| {
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                    async move {
2343                        let value = serde_json::from_str::<serde_json::Value>(&body).unwrap();
2344                        if value.get("method").and_then(serde_json::Value::as_str)
2345                            == Some("initialize")
2346                        {
2347                            return initialize_response().await.into_response();
2348                        }
2349                        if value.get("method").and_then(serde_json::Value::as_str)
2350                            == Some("custom/slow")
2351                        {
2352                            slow_started.notify_one();
2353                            release_slow.notified().await;
2354                        } else if let Some(entries) = value.as_array() {
2355                            if entries.iter().all(|entry| entry.get("method").is_none()) {
2356                                response_batch_seen.notify_one();
2357                            } else {
2358                                call_batch_seen.notify_one();
2359                            }
2360                        }
2361                        StatusCode::ACCEPTED.into_response()
2362                    }
2363                }
2364            })
2365            .get(pending_sse)
2366            .delete(|| async { StatusCode::ACCEPTED }),
2367        );
2368        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2369        let addr = listener.local_addr().unwrap();
2370        let server = tokio::spawn(async move {
2371            axum::serve(listener, app).await.unwrap();
2372        });
2373        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2374        let (mut caller, transport) = Channel::duplex();
2375        let transport = tokio::spawn(run(client, transport));
2376
2377        caller
2378            .tx
2379            .unbounded_send(single_frame(
2380                RawJsonRpcMessage::request(
2381                    "initialize".to_string(),
2382                    json!({}),
2383                    RequestId::Number(1),
2384                )
2385                .unwrap(),
2386            ))
2387            .unwrap();
2388        timeout(Duration::from_secs(1), caller.rx.next())
2389            .await
2390            .unwrap()
2391            .unwrap();
2392
2393        caller
2394            .tx
2395            .unbounded_send(single_frame(
2396                RawJsonRpcMessage::notification("custom/slow".to_string(), json!({})).unwrap(),
2397            ))
2398            .unwrap();
2399        timeout(Duration::from_secs(1), slow_started.notified())
2400            .await
2401            .unwrap();
2402
2403        caller
2404            .tx
2405            .unbounded_send(TransportFrame::Batch(
2406                TransportBatch::from_messages([
2407                    RawJsonRpcMessage::notification("custom/one".to_string(), json!({})).unwrap(),
2408                    RawJsonRpcMessage::notification("custom/two".to_string(), json!({})).unwrap(),
2409                ])
2410                .unwrap(),
2411            ))
2412            .unwrap();
2413        assert!(
2414            timeout(Duration::from_millis(100), call_batch_seen.notified())
2415                .await
2416                .is_err(),
2417            "call-bearing batches must remain behind an earlier ordered POST"
2418        );
2419
2420        caller
2421            .tx
2422            .unbounded_send(TransportFrame::Batch(
2423                TransportBatch::from_messages([
2424                    RawJsonRpcMessage::response(RequestId::Number(10), Ok(json!({}))),
2425                    RawJsonRpcMessage::response(RequestId::Number(11), Ok(json!({}))),
2426                ])
2427                .unwrap(),
2428            ))
2429            .unwrap();
2430        timeout(Duration::from_secs(1), response_batch_seen.notified())
2431            .await
2432            .expect("response-only batch should bypass the ordered POST queue");
2433
2434        release_slow.notify_one();
2435        timeout(Duration::from_secs(1), call_batch_seen.notified())
2436            .await
2437            .expect("call-bearing batch should be sent after the earlier POST completes");
2438
2439        drop(caller);
2440        timeout(Duration::from_secs(1), transport)
2441            .await
2442            .unwrap()
2443            .unwrap()
2444            .unwrap();
2445
2446        server.abort();
2447    }
2448
2449    #[tokio::test]
2450    async fn client_completion_drains_ordered_posts_in_order() {
2451        let first_started = Arc::new(Notify::new());
2452        let release_first = Arc::new(Notify::new());
2453        let second_seen = Arc::new(Notify::new());
2454        let finish_client = Arc::new(Notify::new());
2455        let client_finished = Arc::new(Notify::new());
2456        let (escaped_tx, escaped_rx) = futures::channel::oneshot::channel();
2457        let app = Router::new().route(
2458            "/acp",
2459            post({
2460                let first_started = first_started.clone();
2461                let release_first = release_first.clone();
2462                let second_seen = second_seen.clone();
2463                move |Json(message): Json<RawJsonRpcMessage>| {
2464                    let first_started = first_started.clone();
2465                    let release_first = release_first.clone();
2466                    let second_seen = second_seen.clone();
2467                    async move {
2468                        if is_initialize_request(&message) {
2469                            return initialize_response().await.into_response();
2470                        }
2471
2472                        match method_for_message(&message) {
2473                            Some("custom/first") => {
2474                                first_started.notify_one();
2475                                release_first.notified().await;
2476                            }
2477                            Some("custom/second") => {
2478                                second_seen.notify_one();
2479                            }
2480                            _ => {}
2481                        }
2482                        StatusCode::ACCEPTED.into_response()
2483                    }
2484                }
2485            })
2486            .get(pending_sse)
2487            .delete(|| async { StatusCode::ACCEPTED }),
2488        );
2489        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2490        let addr = listener.local_addr().unwrap();
2491        let server = tokio::spawn(async move {
2492            axum::serve(listener, app).await.unwrap();
2493        });
2494        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2495        let mut connection = tokio::spawn(client.connect_to(PostsThenExitClient {
2496            finish: finish_client.clone(),
2497            finished: client_finished.clone(),
2498            escaped_tx,
2499        }));
2500        let escaped = timeout(Duration::from_secs(1), escaped_rx)
2501            .await
2502            .unwrap()
2503            .unwrap();
2504
2505        timeout(Duration::from_secs(1), first_started.notified())
2506            .await
2507            .unwrap();
2508        assert!(
2509            timeout(Duration::from_millis(100), second_seen.notified())
2510                .await
2511                .is_err(),
2512            "second POST must not be sent while the first POST is pending"
2513        );
2514
2515        finish_client.notify_one();
2516        timeout(Duration::from_secs(1), client_finished.notified())
2517            .await
2518            .unwrap();
2519        assert!(
2520            timeout(Duration::from_millis(100), &mut connection)
2521                .await
2522                .is_err(),
2523            "HTTP transport returned before its accepted POSTs completed"
2524        );
2525        assert!(
2526            escaped
2527                .unbounded_send(single_frame(
2528                    RawJsonRpcMessage::notification("custom/too-late".to_string(), json!({}),)
2529                        .unwrap()
2530                ))
2531                .is_err(),
2532            "escaped client sender remained open after client completion"
2533        );
2534
2535        release_first.notify_one();
2536        timeout(Duration::from_secs(1), second_seen.notified())
2537            .await
2538            .unwrap();
2539
2540        timeout(Duration::from_secs(1), connection)
2541            .await
2542            .unwrap()
2543            .unwrap()
2544            .unwrap();
2545
2546        server.abort();
2547    }
2548
2549    #[tokio::test]
2550    async fn client_completion_cancels_pending_sse_establishment() {
2551        let sse_started = Arc::new(Notify::new());
2552        let delete_count = Arc::new(AtomicUsize::new(0));
2553        let client_finished = Arc::new(Notify::new());
2554        let app = Router::new().route(
2555            "/acp",
2556            post(initialize_response)
2557                .get({
2558                    let sse_started = sse_started.clone();
2559                    move || {
2560                        let sse_started = sse_started.clone();
2561                        async move {
2562                            sse_started.notify_one();
2563                            futures::future::pending::<StatusCode>().await
2564                        }
2565                    }
2566                })
2567                .delete({
2568                    let delete_count = delete_count.clone();
2569                    move || {
2570                        let delete_count = delete_count.clone();
2571                        async move {
2572                            delete_count.fetch_add(1, Ordering::SeqCst);
2573                            StatusCode::ACCEPTED
2574                        }
2575                    }
2576                }),
2577        );
2578        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2579        let addr = listener.local_addr().unwrap();
2580        let server = tokio::spawn(async move {
2581            axum::serve(listener, app).await.unwrap();
2582        });
2583        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2584        let connection = tokio::spawn(client.connect_to(InitializeThenExitClient {
2585            sse_started,
2586            finished: client_finished.clone(),
2587        }));
2588
2589        timeout(Duration::from_secs(1), client_finished.notified())
2590            .await
2591            .expect("client foreground did not finish after the SSE request started");
2592
2593        timeout(Duration::from_secs(1), connection)
2594            .await
2595            .expect("transport remained blocked on SSE response headers")
2596            .unwrap()
2597            .unwrap();
2598        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
2599
2600        server.abort();
2601    }
2602
2603    #[tokio::test]
2604    async fn stalled_sse_establishment_observes_earlier_post_failure() {
2605        let app = Router::new().route(
2606            "/acp",
2607            get(|| async { futures::future::pending::<StatusCode>().await }),
2608        );
2609        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2610        let addr = listener.local_addr().unwrap();
2611        let server = tokio::spawn(async move {
2612            axum::serve(listener, app).await.unwrap();
2613        });
2614
2615        let connection = HttpConnection::new(
2616            url::Url::parse(&format!("http://{addr}/acp")).unwrap(),
2617            reqwest::Client::new(),
2618        );
2619        connection.set_connection_id("connection-1".to_string());
2620        let (incoming, _incoming_rx) = mpsc::unbounded();
2621        let mut state = ClientState {
2622            connection: connection.clone(),
2623            open_session_streams: HashSet::new(),
2624            pending_requests: HashMap::new(),
2625            incoming,
2626        };
2627        let pending_request = (RequestId::Number(7), "custom/earlier".to_string());
2628        state.track_pending_requests(std::slice::from_ref(&pending_request));
2629        let mut posts = PostQueues::default();
2630        posts.ordered.push(PendingPost {
2631            pending_requests: vec![pending_request],
2632            response: async { Err("earlier post failed".to_string()) }.boxed(),
2633        });
2634
2635        let (_outgoing_tx, mut outgoing) = mpsc::unbounded();
2636        let mut buffered_outgoing = VecDeque::new();
2637        let (event_tx, mut event_rx) = mpsc::unbounded();
2638        let mut lifecycle = HttpTransportLifecycle::new(connection);
2639        let error = timeout(
2640            Duration::from_secs(1),
2641            lifecycle.start_sse(
2642                Some("later-session".to_string()),
2643                event_tx,
2644                SseStartContext {
2645                    events: &mut event_rx,
2646                    outgoing: &mut outgoing,
2647                    buffered_outgoing: &mut buffered_outgoing,
2648                    posts: &mut posts,
2649                    state: &mut state,
2650                },
2651            ),
2652        )
2653        .await
2654        .expect("stalled SSE setup hid an earlier POST failure")
2655        .unwrap_err();
2656
2657        assert!(error.to_string().contains("earlier post failed"));
2658        assert!(state.pending_requests.is_empty());
2659
2660        lifecycle.close().await;
2661        server.abort();
2662    }
2663
2664    #[tokio::test]
2665    async fn stalled_sse_establishment_keeps_callback_responses_moving() {
2666        let release_get = Arc::new(Notify::new());
2667        let complete_earlier_post = Arc::new(Notify::new());
2668        let app = Router::new().route(
2669            "/acp",
2670            post({
2671                let release_get = release_get.clone();
2672                let complete_earlier_post = complete_earlier_post.clone();
2673                move || {
2674                    release_get.notify_one();
2675                    complete_earlier_post.notify_one();
2676                    async { StatusCode::ACCEPTED }
2677                }
2678            })
2679            .get({
2680                let release_get = release_get.clone();
2681                move || {
2682                    let release_get = release_get.clone();
2683                    async move {
2684                        release_get.notified().await;
2685                        Sse::new(futures::stream::pending::<Result<Event, Infallible>>())
2686                    }
2687                }
2688            }),
2689        );
2690        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2691        let addr = listener.local_addr().unwrap();
2692        let server = tokio::spawn(async move {
2693            axum::serve(listener, app).await.unwrap();
2694        });
2695
2696        let connection = HttpConnection::new(
2697            url::Url::parse(&format!("http://{addr}/acp")).unwrap(),
2698            reqwest::Client::new(),
2699        );
2700        connection.set_connection_id("connection-1".to_string());
2701        let (incoming, mut incoming_rx) = mpsc::unbounded();
2702        let mut state = ClientState {
2703            connection: connection.clone(),
2704            open_session_streams: HashSet::new(),
2705            pending_requests: HashMap::new(),
2706            incoming,
2707        };
2708        let mut posts = PostQueues::default();
2709        posts.ordered.push(PendingPost {
2710            pending_requests: Vec::new(),
2711            response: async move {
2712                complete_earlier_post.notified().await;
2713                Ok(())
2714            }
2715            .boxed(),
2716        });
2717
2718        let (outgoing_tx, mut outgoing) = mpsc::unbounded();
2719        let outgoing_guard = outgoing_tx.clone();
2720        let mut buffered_outgoing = VecDeque::new();
2721        let (event_tx, mut event_rx) = mpsc::unbounded();
2722        event_tx
2723            .unbounded_send(SseMessage {
2724                frame: single_frame(
2725                    RawJsonRpcMessage::request(
2726                        "test/callback".to_string(),
2727                        json!({}),
2728                        RequestId::Number(99),
2729                    )
2730                    .unwrap(),
2731                ),
2732            })
2733            .unwrap();
2734
2735        let responder = async move {
2736            let callback = incoming_rx
2737                .next()
2738                .await
2739                .expect("callback was not delivered");
2740            assert!(matches!(
2741                into_single_message(callback).unwrap(),
2742                RawJsonRpcMessage::Request(request)
2743                    if request.method.as_ref() == "test/callback"
2744            ));
2745            outgoing_tx
2746                .unbounded_send(single_frame(RawJsonRpcMessage::response(
2747                    RequestId::Number(99),
2748                    Ok(json!({})),
2749                )))
2750                .unwrap();
2751        };
2752        let mut lifecycle = HttpTransportLifecycle::new(connection);
2753        let (outcome, ()) = timeout(Duration::from_secs(1), async {
2754            futures::join!(
2755                lifecycle.start_sse(
2756                    Some("later-session".to_string()),
2757                    event_tx,
2758                    SseStartContext {
2759                        events: &mut event_rx,
2760                        outgoing: &mut outgoing,
2761                        buffered_outgoing: &mut buffered_outgoing,
2762                        posts: &mut posts,
2763                        state: &mut state,
2764                    },
2765                ),
2766                responder,
2767            )
2768        })
2769        .await
2770        .expect("callback response deadlocked behind stalled SSE establishment");
2771
2772        assert_eq!(outcome.unwrap(), SseStartOutcome::Established);
2773        assert!(buffered_outgoing.is_empty());
2774
2775        drop(outgoing_guard);
2776        lifecycle.close().await;
2777        server.abort();
2778    }
2779
2780    #[tokio::test]
2781    async fn pending_sse_establishment_reports_buffered_output_on_shutdown() {
2782        let sse_started = Arc::new(Notify::new());
2783        let delete_count = Arc::new(AtomicUsize::new(0));
2784        let app = Router::new().route(
2785            "/acp",
2786            post(initialize_response)
2787                .get({
2788                    let sse_started = sse_started.clone();
2789                    move || {
2790                        let sse_started = sse_started.clone();
2791                        async move {
2792                            sse_started.notify_one();
2793                            futures::future::pending::<StatusCode>().await
2794                        }
2795                    }
2796                })
2797                .delete({
2798                    let delete_count = delete_count.clone();
2799                    move || {
2800                        let delete_count = delete_count.clone();
2801                        async move {
2802                            delete_count.fetch_add(1, Ordering::SeqCst);
2803                            StatusCode::ACCEPTED
2804                        }
2805                    }
2806                }),
2807        );
2808        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2809        let addr = listener.local_addr().unwrap();
2810        let server = tokio::spawn(async move {
2811            axum::serve(listener, app).await.unwrap();
2812        });
2813        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2814        let (mut caller, transport) = Channel::duplex();
2815        let transport = tokio::spawn(run(client, transport));
2816
2817        caller
2818            .tx
2819            .unbounded_send(single_frame(
2820                RawJsonRpcMessage::request(
2821                    "initialize".to_string(),
2822                    json!({}),
2823                    RequestId::Number(1),
2824                )
2825                .unwrap(),
2826            ))
2827            .unwrap();
2828        timeout(Duration::from_secs(1), caller.rx.next())
2829            .await
2830            .unwrap()
2831            .unwrap();
2832        timeout(Duration::from_secs(1), sse_started.notified())
2833            .await
2834            .expect("connection SSE request did not reach the server");
2835
2836        caller
2837            .tx
2838            .unbounded_send(single_frame(
2839                RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(),
2840            ))
2841            .unwrap();
2842        drop(caller);
2843
2844        let error = timeout(Duration::from_secs(1), transport)
2845            .await
2846            .expect("transport remained blocked on SSE response headers")
2847            .unwrap()
2848            .unwrap_err();
2849        assert!(error.to_string().contains("accepted messages"));
2850        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
2851
2852        server.abort();
2853    }
2854
2855    #[tokio::test]
2856    async fn sse_continues_while_post_is_pending() {
2857        let post_started = Arc::new(Notify::new());
2858        let callback_response_seen = Arc::new(Notify::new());
2859        let sse_started = Arc::new(Notify::new());
2860        let (callback_tx, mut callback_rx) = tokio::sync::mpsc::unbounded_channel();
2861        let app = Router::new().route(
2862            "/acp",
2863            post({
2864                let post_started = post_started.clone();
2865                let callback_response_seen = callback_response_seen.clone();
2866                let callback_tx = callback_tx.clone();
2867                move |Json(message): Json<RawJsonRpcMessage>| {
2868                    let post_started = post_started.clone();
2869                    let callback_response_seen = callback_response_seen.clone();
2870                    let callback_tx = callback_tx.clone();
2871                    async move {
2872                        if is_initialize_request(&message) {
2873                            return initialize_response().await.into_response();
2874                        }
2875
2876                        match &message {
2877                            RawJsonRpcMessage::Request(request)
2878                                if request.method.as_ref() == "custom/slow" =>
2879                            {
2880                                post_started.notify_waiters();
2881                                callback_response_seen.notified().await;
2882                                StatusCode::ACCEPTED.into_response()
2883                            }
2884                            RawJsonRpcMessage::Response(
2885                                RpcResponse::Result {
2886                                    id: RequestId::Number(99),
2887                                    ..
2888                                }
2889                                | RpcResponse::Error {
2890                                    id: RequestId::Number(99),
2891                                    ..
2892                                },
2893                            ) => {
2894                                callback_tx.send(message).unwrap();
2895                                callback_response_seen.notify_waiters();
2896                                StatusCode::ACCEPTED.into_response()
2897                            }
2898                            _ => StatusCode::ACCEPTED.into_response(),
2899                        }
2900                    }
2901                }
2902            })
2903            .get({
2904                let post_started = post_started.clone();
2905                let sse_started = sse_started.clone();
2906                move || {
2907                    let post_started = post_started.clone();
2908                    let sse_started = sse_started.clone();
2909                    async move {
2910                        let stream = async_stream::stream! {
2911                            sse_started.notify_waiters();
2912                            post_started.notified().await;
2913                            yield Ok::<_, Infallible>(sse_event(
2914                                RawJsonRpcMessage::request(
2915                                    "client/callback".to_string(),
2916                                    json!({}),
2917                                    RequestId::Number(99),
2918                                )
2919                                .unwrap(),
2920                            ));
2921                            futures::future::pending::<()>().await;
2922                        };
2923                        Sse::new(stream)
2924                    }
2925                }
2926            })
2927            .delete(|| async { StatusCode::ACCEPTED }),
2928        );
2929        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2930        let addr = listener.local_addr().unwrap();
2931        let server = tokio::spawn(async move {
2932            axum::serve(listener, app).await.unwrap();
2933        });
2934        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2935        let (mut caller, transport) = Channel::duplex();
2936        let transport = tokio::spawn(run(client, transport));
2937
2938        caller
2939            .tx
2940            .unbounded_send(single_frame(
2941                RawJsonRpcMessage::request(
2942                    "initialize".to_string(),
2943                    json!({}),
2944                    RequestId::Number(1),
2945                )
2946                .unwrap(),
2947            ))
2948            .unwrap();
2949        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2950            .await
2951            .unwrap()
2952            .unwrap()
2953            .unwrap();
2954        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2955        timeout(Duration::from_secs(1), sse_started.notified())
2956            .await
2957            .unwrap();
2958
2959        caller
2960            .tx
2961            .unbounded_send(single_frame(
2962                RawJsonRpcMessage::request(
2963                    "custom/slow".to_string(),
2964                    json!({}),
2965                    RequestId::Number(2),
2966                )
2967                .unwrap(),
2968            ))
2969            .unwrap();
2970
2971        let callback = timeout(Duration::from_secs(1), caller.rx.next())
2972            .await
2973            .unwrap()
2974            .unwrap()
2975            .unwrap();
2976        assert!(matches!(
2977            callback,
2978            RawJsonRpcMessage::Request(request)
2979                if request.method.as_ref() == "client/callback"
2980                    && request.id == RequestId::Number(99)
2981        ));
2982
2983        caller
2984            .tx
2985            .unbounded_send(single_frame(RawJsonRpcMessage::response(
2986                RequestId::Number(99),
2987                Ok(json!({})),
2988            )))
2989            .unwrap();
2990        let callback_response = timeout(Duration::from_secs(1), callback_rx.recv())
2991            .await
2992            .unwrap()
2993            .unwrap();
2994        assert!(matches!(
2995            callback_response,
2996            RawJsonRpcMessage::Response(RpcResponse::Result {
2997                id: RequestId::Number(99),
2998                ..
2999            })
3000        ));
3001
3002        drop(caller);
3003        timeout(Duration::from_secs(1), transport)
3004            .await
3005            .unwrap()
3006            .unwrap()
3007            .unwrap();
3008
3009        server.abort();
3010    }
3011
3012    #[tokio::test]
3013    async fn post_error_deletes_initialized_connection() {
3014        let delete_count = Arc::new(AtomicUsize::new(0));
3015        let delete_count_for_handler = delete_count.clone();
3016        let app = Router::new().route(
3017            "/acp",
3018            post(initialize_response).get(pending_sse).delete(move || {
3019                let delete_count = delete_count_for_handler.clone();
3020                async move {
3021                    delete_count.fetch_add(1, Ordering::SeqCst);
3022                    StatusCode::ACCEPTED
3023                }
3024            }),
3025        );
3026        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3027        let addr = listener.local_addr().unwrap();
3028        let server = tokio::spawn(async move {
3029            axum::serve(listener, app).await.unwrap();
3030        });
3031        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3032        let (mut caller, transport) = Channel::duplex();
3033        let transport = tokio::spawn(run(client, transport));
3034
3035        caller
3036            .tx
3037            .unbounded_send(single_frame(
3038                RawJsonRpcMessage::request(
3039                    "initialize".to_string(),
3040                    json!({}),
3041                    RequestId::Number(1),
3042                )
3043                .unwrap(),
3044            ))
3045            .unwrap();
3046        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3047            .await
3048            .unwrap()
3049            .unwrap()
3050            .unwrap();
3051        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3052
3053        caller
3054            .tx
3055            .unbounded_send(single_frame(
3056                RawJsonRpcMessage::request(
3057                    "session/prompt".to_string(),
3058                    json!({}),
3059                    RequestId::Number(2),
3060                )
3061                .unwrap(),
3062            ))
3063            .unwrap();
3064        let error = timeout(Duration::from_secs(1), transport)
3065            .await
3066            .unwrap()
3067            .unwrap()
3068            .unwrap_err();
3069
3070        assert!(error.to_string().contains("POST"));
3071        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3072
3073        server.abort();
3074    }
3075
3076    #[tokio::test]
3077    async fn connection_sse_disconnect_fails_transport() {
3078        let delete_count = Arc::new(AtomicUsize::new(0));
3079        let delete_count_for_handler = delete_count.clone();
3080        let app = Router::new().route(
3081            "/acp",
3082            post(initialize_response).get(closed_sse).delete(move || {
3083                let delete_count = delete_count_for_handler.clone();
3084                async move {
3085                    delete_count.fetch_add(1, Ordering::SeqCst);
3086                    StatusCode::ACCEPTED
3087                }
3088            }),
3089        );
3090        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3091        let addr = listener.local_addr().unwrap();
3092        let server = tokio::spawn(async move {
3093            axum::serve(listener, app).await.unwrap();
3094        });
3095        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3096        let (mut caller, transport) = Channel::duplex();
3097        let transport = tokio::spawn(run(client, transport));
3098
3099        caller
3100            .tx
3101            .unbounded_send(single_frame(
3102                RawJsonRpcMessage::request(
3103                    "initialize".to_string(),
3104                    json!({}),
3105                    RequestId::Number(1),
3106                )
3107                .unwrap(),
3108            ))
3109            .unwrap();
3110        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3111            .await
3112            .unwrap()
3113            .unwrap()
3114            .unwrap();
3115        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3116
3117        let error = timeout(Duration::from_secs(1), transport)
3118            .await
3119            .unwrap()
3120            .unwrap()
3121            .unwrap_err();
3122
3123        assert!(error.to_string().contains("SSE"));
3124        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3125
3126        server.abort();
3127    }
3128
3129    #[tokio::test]
3130    async fn malformed_sse_json_is_delivered_and_transport_continues() {
3131        let delete_count = Arc::new(AtomicUsize::new(0));
3132        let delete_count_for_handler = delete_count.clone();
3133        let app = Router::new().route(
3134            "/acp",
3135            post(initialize_response)
3136                .get(malformed_sse)
3137                .delete(move || {
3138                    let delete_count = delete_count_for_handler.clone();
3139                    async move {
3140                        delete_count.fetch_add(1, Ordering::SeqCst);
3141                        StatusCode::ACCEPTED
3142                    }
3143                }),
3144        );
3145        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3146        let addr = listener.local_addr().unwrap();
3147        let server = tokio::spawn(async move {
3148            axum::serve(listener, app).await.unwrap();
3149        });
3150        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3151        let (mut caller, transport) = Channel::duplex();
3152        let transport = tokio::spawn(run(client, transport));
3153
3154        caller
3155            .tx
3156            .unbounded_send(single_frame(
3157                RawJsonRpcMessage::request(
3158                    "initialize".to_string(),
3159                    json!({}),
3160                    RequestId::Number(1),
3161                )
3162                .unwrap(),
3163            ))
3164            .unwrap();
3165        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3166            .await
3167            .unwrap()
3168            .unwrap()
3169            .unwrap();
3170        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3171
3172        let frame = timeout(Duration::from_secs(1), caller.rx.next())
3173            .await
3174            .unwrap()
3175            .unwrap();
3176
3177        let TransportFrame::Malformed { raw, error } = frame else {
3178            panic!("expected malformed frame, got {frame:?}");
3179        };
3180        assert_eq!(raw, "{not json");
3181        assert_eq!(error.code, AcpError::parse_error().code);
3182        drop(caller);
3183        timeout(Duration::from_secs(1), transport)
3184            .await
3185            .unwrap()
3186            .unwrap()
3187            .unwrap();
3188        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3189
3190        server.abort();
3191    }
3192
3193    #[tokio::test]
3194    async fn malformed_ws_json_reports_parse_error_and_continues() {
3195        let app = Router::new().route("/acp", get(malformed_then_valid_ws));
3196        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3197        let addr = listener.local_addr().unwrap();
3198        let server = tokio::spawn(async move {
3199            axum::serve(listener, app).await.unwrap();
3200        });
3201        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
3202        let (mut caller, transport) = Channel::duplex();
3203        let transport = tokio::spawn(run(client, transport));
3204
3205        let frame = timeout(Duration::from_secs(1), caller.rx.next())
3206            .await
3207            .unwrap()
3208            .unwrap();
3209        let TransportFrame::Malformed { raw, error } = frame else {
3210            panic!("expected malformed frame, got {frame:?}");
3211        };
3212        assert_eq!(raw, "{not json");
3213        assert_eq!(error.code, AcpError::parse_error().code);
3214
3215        let message = timeout(Duration::from_secs(1), caller.rx.next())
3216            .await
3217            .unwrap()
3218            .unwrap()
3219            .unwrap();
3220        assert!(matches!(message, RawJsonRpcMessage::Response(_)));
3221
3222        drop(caller);
3223        timeout(Duration::from_secs(1), transport)
3224            .await
3225            .unwrap()
3226            .unwrap()
3227            .unwrap();
3228
3229        server.abort();
3230    }
3231
3232    #[tokio::test]
3233    async fn websocket_serializes_batch_as_one_text_frame() {
3234        let (caller, transport) = Channel::duplex();
3235        let Channel {
3236            tx: outgoing,
3237            rx: incoming,
3238        } = caller;
3239        drop(incoming);
3240        outgoing
3241            .unbounded_send(TransportFrame::Batch(
3242                TransportBatch::from_messages([
3243                    RawJsonRpcMessage::notification("custom/first".to_string(), json!({})).unwrap(),
3244                    RawJsonRpcMessage::notification("custom/second".to_string(), json!({}))
3245                        .unwrap(),
3246                ])
3247                .unwrap(),
3248            ))
3249            .unwrap();
3250        drop(outgoing);
3251
3252        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
3253        timeout(
3254            Duration::from_secs(1),
3255            drive_ws(
3256                RecordingWsSink(ws_output_tx),
3257                futures::stream::pending::<Result<WsMessage, std::io::Error>>(),
3258                transport,
3259            ),
3260        )
3261        .await
3262        .unwrap()
3263        .unwrap();
3264        let frames = ws_output.by_ref().collect::<Vec<_>>().await;
3265
3266        let WsMessage::Text(text) = &frames[0] else {
3267            panic!("batch was not sent as WebSocket text");
3268        };
3269        let batch = serde_json::from_str::<serde_json::Value>(text.as_str()).unwrap();
3270        let entries = batch.as_array().expect("batch should remain an array");
3271        assert_eq!(entries.len(), 2);
3272        assert_eq!(entries[0]["method"], "custom/first");
3273        assert_eq!(entries[1]["method"], "custom/second");
3274        assert!(matches!(frames.get(1), Some(WsMessage::Close(None))));
3275        assert_eq!(frames.len(), 2);
3276    }
3277
3278    #[tokio::test]
3279    async fn websocket_drain_discards_incoming_after_receiver_closes() {
3280        let (caller, transport) = Channel::duplex();
3281        let Channel {
3282            tx: outgoing,
3283            rx: incoming,
3284        } = caller;
3285        drop(incoming);
3286
3287        let inbound =
3288            RawJsonRpcMessage::notification("custom/inbound".to_string(), json!({})).unwrap();
3289        let inbound = WsMessage::Text(serde_json::to_string(&inbound).unwrap().into());
3290        let ws_rx = QueueOutgoingThenText {
3291            text: Some(inbound),
3292            outgoing: Some(outgoing),
3293        };
3294        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
3295        timeout(
3296            Duration::from_secs(1),
3297            drive_ws(RecordingWsSink(ws_output_tx), ws_rx, transport),
3298        )
3299        .await
3300        .unwrap()
3301        .unwrap();
3302        let mut frames = Vec::new();
3303        while let Some(frame) = ws_output.next().await {
3304            frames.push(frame);
3305        }
3306
3307        let messages = frames
3308            .iter()
3309            .filter_map(|frame| match frame {
3310                WsMessage::Text(text) => {
3311                    Some(serde_json::from_str::<RawJsonRpcMessage>(text.as_str()).unwrap())
3312                }
3313                _ => None,
3314            })
3315            .collect::<Vec<_>>();
3316        let methods = messages
3317            .iter()
3318            .filter_map(method_for_message)
3319            .collect::<Vec<_>>();
3320        assert_eq!(methods, ["custom/first", "custom/second"]);
3321        assert!(matches!(frames.last(), Some(WsMessage::Close(None))));
3322    }
3323
3324    #[tokio::test]
3325    async fn websocket_reader_runs_while_send_is_backpressured() {
3326        let (caller, transport) = Channel::duplex();
3327        let Channel {
3328            tx: outgoing,
3329            rx: incoming,
3330        } = caller;
3331        drop(incoming);
3332        outgoing
3333            .unbounded_send(single_frame(
3334                RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(),
3335            ))
3336            .unwrap();
3337        drop(outgoing);
3338
3339        let (started_tx, started_rx) = mpsc::unbounded();
3340        let (release_tx, release_rx) = futures::channel::oneshot::channel();
3341        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
3342        let ws_tx = BackpressuredWsSink {
3343            output: ws_output_tx,
3344            started: started_tx,
3345            release: Some(release_rx),
3346        };
3347        let ws_rx = ReleaseBackpressureOnPoll {
3348            started: started_rx,
3349            release: Some(release_tx),
3350        };
3351
3352        timeout(Duration::from_secs(1), drive_ws(ws_tx, ws_rx, transport))
3353            .await
3354            .expect("WebSocket reader was not polled while its writer was backpressured")
3355            .unwrap();
3356        let frames = ws_output.by_ref().collect::<Vec<_>>().await;
3357
3358        let WsMessage::Text(text) = &frames[0] else {
3359            panic!("queued message was not sent as WebSocket text");
3360        };
3361        let message = serde_json::from_str::<RawJsonRpcMessage>(text.as_str()).unwrap();
3362        assert_eq!(method_for_message(&message), Some("custom/queued"));
3363        assert!(matches!(frames.get(1), Some(WsMessage::Close(None))));
3364        assert_eq!(frames.len(), 2);
3365    }
3366
3367    #[tokio::test]
3368    async fn peer_ws_close_fails_transport() {
3369        let app = Router::new().route("/acp", get(close_ws));
3370        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3371        let addr = listener.local_addr().unwrap();
3372        let server = tokio::spawn(async move {
3373            axum::serve(listener, app).await.unwrap();
3374        });
3375        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
3376        let (_caller, transport) = Channel::duplex();
3377        let transport = tokio::spawn(run(client, transport));
3378
3379        let error = timeout(Duration::from_secs(1), transport)
3380            .await
3381            .unwrap()
3382            .unwrap()
3383            .unwrap_err();
3384        assert!(error.to_string().contains("WebSocket closed by peer"));
3385
3386        server.abort();
3387    }
3388
3389    #[tokio::test]
3390    async fn dropped_transport_future_deletes_initialized_connection() {
3391        let delete_count = Arc::new(AtomicUsize::new(0));
3392        let delete_count_for_handler = delete_count.clone();
3393        let app = Router::new().route(
3394            "/acp",
3395            post(initialize_response).get(pending_sse).delete(move || {
3396                let delete_count = delete_count_for_handler.clone();
3397                async move {
3398                    delete_count.fetch_add(1, Ordering::SeqCst);
3399                    StatusCode::ACCEPTED
3400                }
3401            }),
3402        );
3403        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3404        let addr = listener.local_addr().unwrap();
3405        let server = tokio::spawn(async move {
3406            axum::serve(listener, app).await.unwrap();
3407        });
3408        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3409        let (mut caller, transport) = Channel::duplex();
3410        let mut transport = Box::pin(run(client, transport));
3411
3412        caller
3413            .tx
3414            .unbounded_send(single_frame(
3415                RawJsonRpcMessage::request(
3416                    "initialize".to_string(),
3417                    json!({}),
3418                    RequestId::Number(1),
3419                )
3420                .unwrap(),
3421            ))
3422            .unwrap();
3423        let init_response = timeout(Duration::from_secs(1), async {
3424            tokio::select! {
3425                result = &mut transport => {
3426                    panic!("transport ended before initialize response: {result:?}");
3427                }
3428                msg = caller.rx.next() => {
3429                    msg.unwrap().unwrap()
3430                }
3431            }
3432        })
3433        .await
3434        .unwrap();
3435        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3436
3437        drop(transport);
3438        wait_for_delete(&delete_count).await;
3439
3440        server.abort();
3441    }
3442
3443    #[tokio::test]
3444    async fn dropped_transport_during_close_retries_delete() {
3445        let delete_count = Arc::new(AtomicUsize::new(0));
3446        let delete_count_for_handler = delete_count.clone();
3447        let release_delete = Arc::new(Notify::new());
3448        let release_delete_for_handler = release_delete.clone();
3449        let app = Router::new().route(
3450            "/acp",
3451            post(initialize_response).get(pending_sse).delete(move || {
3452                let delete_count = delete_count_for_handler.clone();
3453                let release_delete = release_delete_for_handler.clone();
3454                async move {
3455                    delete_count.fetch_add(1, Ordering::SeqCst);
3456                    release_delete.notified().await;
3457                    StatusCode::ACCEPTED
3458                }
3459            }),
3460        );
3461        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3462        let addr = listener.local_addr().unwrap();
3463        let server = tokio::spawn(async move {
3464            axum::serve(listener, app).await.unwrap();
3465        });
3466        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3467        let (mut caller, transport) = Channel::duplex();
3468        let transport = tokio::spawn(run(client, transport));
3469
3470        caller
3471            .tx
3472            .unbounded_send(single_frame(
3473                RawJsonRpcMessage::request(
3474                    "initialize".to_string(),
3475                    json!({}),
3476                    RequestId::Number(1),
3477                )
3478                .unwrap(),
3479            ))
3480            .unwrap();
3481        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3482            .await
3483            .unwrap()
3484            .unwrap()
3485            .unwrap();
3486        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3487
3488        drop(caller);
3489        wait_for_delete_count(&delete_count, 1).await;
3490        transport.abort();
3491        wait_for_delete_count(&delete_count, 2).await;
3492        release_delete.notify_waiters();
3493        drop(transport.await);
3494
3495        server.abort();
3496    }
3497
3498    #[tokio::test]
3499    async fn initialize_error_without_connection_id_is_delivered_without_sse() {
3500        let get_count = Arc::new(AtomicUsize::new(0));
3501        let get_count_for_handler = get_count.clone();
3502        let app = Router::new().route(
3503            "/acp",
3504            post(initialize_error_response).get(move || {
3505                let get_count = get_count_for_handler.clone();
3506                async move {
3507                    get_count.fetch_add(1, Ordering::SeqCst);
3508                    pending_sse().await
3509                }
3510            }),
3511        );
3512        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3513        let addr = listener.local_addr().unwrap();
3514        let server = tokio::spawn(async move {
3515            axum::serve(listener, app).await.unwrap();
3516        });
3517        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3518        let (mut caller, transport) = Channel::duplex();
3519        let transport = tokio::spawn(run(client, transport));
3520
3521        caller
3522            .tx
3523            .unbounded_send(single_frame(
3524                RawJsonRpcMessage::request(
3525                    "initialize".to_string(),
3526                    json!({}),
3527                    RequestId::Number(1),
3528                )
3529                .unwrap(),
3530            ))
3531            .unwrap();
3532        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3533            .await
3534            .unwrap()
3535            .unwrap()
3536            .unwrap();
3537
3538        assert!(matches!(
3539            init_response,
3540            RawJsonRpcMessage::Response(RpcResponse::Error {
3541                id: RequestId::Number(1),
3542                ..
3543            })
3544        ));
3545        assert_eq!(get_count.load(Ordering::SeqCst), 0);
3546
3547        drop(caller);
3548        timeout(Duration::from_secs(1), transport)
3549            .await
3550            .unwrap()
3551            .unwrap()
3552            .unwrap();
3553
3554        server.abort();
3555    }
3556
3557    #[tokio::test]
3558    async fn malformed_initialize_body_with_connection_id_is_deleted() {
3559        let delete_count = Arc::new(AtomicUsize::new(0));
3560        let delete_count_for_handler = delete_count.clone();
3561        let app = Router::new().route(
3562            "/acp",
3563            post(malformed_initialize_response).delete(move || {
3564                let delete_count = delete_count_for_handler.clone();
3565                async move {
3566                    delete_count.fetch_add(1, Ordering::SeqCst);
3567                    StatusCode::ACCEPTED
3568                }
3569            }),
3570        );
3571        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3572        let addr = listener.local_addr().unwrap();
3573        let server = tokio::spawn(async move {
3574            axum::serve(listener, app).await.unwrap();
3575        });
3576        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3577        let (caller, transport) = Channel::duplex();
3578        let transport = tokio::spawn(run(client, transport));
3579
3580        caller
3581            .tx
3582            .unbounded_send(single_frame(
3583                RawJsonRpcMessage::request(
3584                    "initialize".to_string(),
3585                    json!({}),
3586                    RequestId::Number(1),
3587                )
3588                .unwrap(),
3589            ))
3590            .unwrap();
3591        let error = timeout(Duration::from_secs(1), transport)
3592            .await
3593            .unwrap()
3594            .unwrap()
3595            .unwrap_err();
3596
3597        assert!(error.to_string().contains("initialize"));
3598        wait_for_delete(&delete_count).await;
3599
3600        server.abort();
3601    }
3602
3603    async fn wait_for_delete(delete_count: &AtomicUsize) {
3604        wait_for_delete_count(delete_count, 1).await;
3605        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3606    }
3607
3608    async fn wait_for_delete_count(delete_count: &AtomicUsize, expected: usize) {
3609        timeout(Duration::from_secs(1), async {
3610            loop {
3611                if delete_count.load(Ordering::SeqCst) >= expected {
3612                    break;
3613                }
3614                sleep(Duration::from_millis(10)).await;
3615            }
3616        })
3617        .await
3618        .unwrap();
3619    }
3620
3621    async fn initialize_response() -> impl IntoResponse {
3622        let mut headers = HeaderMap::new();
3623        headers.insert(HEADER_CONNECTION_ID, HeaderValue::from_static("conn-1"));
3624        (
3625            StatusCode::OK,
3626            headers,
3627            Json(RawJsonRpcMessage::response(
3628                RequestId::Number(1),
3629                Ok(json!({})),
3630            )),
3631        )
3632    }
3633
3634    async fn initialize_error_response() -> Json<RawJsonRpcMessage> {
3635        Json(RawJsonRpcMessage::response(
3636            RequestId::Number(1),
3637            Err(AcpError::invalid_request().data("initialize rejected")),
3638        ))
3639    }
3640
3641    async fn malformed_initialize_response() -> impl IntoResponse {
3642        let mut headers = HeaderMap::new();
3643        headers.insert(HEADER_CONNECTION_ID, HeaderValue::from_static("conn-1"));
3644        (StatusCode::OK, headers, "{not json")
3645    }
3646
3647    async fn pending_sse() -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
3648        Sse::new(futures::stream::pending())
3649    }
3650
3651    fn sse_event(message: RawJsonRpcMessage) -> Event {
3652        Event::default().data(serde_json::to_string(&message).unwrap())
3653    }
3654
3655    async fn malformed_sse() -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
3656        let invalid = futures::stream::once(async {
3657            Ok::<_, Infallible>(Event::default().data("{not json"))
3658        });
3659        Sse::new(invalid.chain(futures::stream::pending()))
3660    }
3661
3662    async fn malformed_then_valid_ws(ws: WebSocketUpgrade) -> impl IntoResponse {
3663        ws.on_upgrade(|mut socket| async move {
3664            drop(socket.send(AxumWsMessage::Text("{not json".into())).await);
3665            let valid = serde_json::to_string(&RawJsonRpcMessage::response(
3666                RequestId::Number(1),
3667                Ok(json!({})),
3668            ))
3669            .unwrap();
3670            drop(socket.send(AxumWsMessage::Text(valid.into())).await);
3671            futures::future::pending::<()>().await;
3672        })
3673    }
3674
3675    async fn close_ws(ws: WebSocketUpgrade) -> impl IntoResponse {
3676        ws.on_upgrade(|mut socket| async move {
3677            drop(socket.send(AxumWsMessage::Close(None)).await);
3678        })
3679    }
3680
3681    async fn closed_sse() -> StatusCode {
3682        StatusCode::OK
3683    }
3684}