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,
8    schema::v1::{RequestId, Response as RpcResponse},
9};
10use async_tungstenite::tungstenite::Message as WsMessage;
11use futures::{
12    Stream, StreamExt,
13    channel::mpsc::{self, UnboundedSender},
14    future::{BoxFuture, FutureExt},
15    pin_mut,
16    stream::FuturesUnordered,
17};
18use thiserror::Error;
19use tracing::{debug, error, trace, warn};
20
21use crate::protocol::{
22    HEADER_CONNECTION_ID, HEADER_SESSION_ID, is_initialize_request, method_for_message,
23    method_requires_session_header, session_id_from_message,
24};
25
26#[derive(Debug, Error)]
27pub enum HttpClientError {
28    #[error("invalid URL: {0}")]
29    InvalidUrl(#[from] url::ParseError),
30    #[error("failed to build HTTP client: {0}")]
31    Reqwest(#[from] reqwest::Error),
32}
33
34pub struct HttpClient {
35    endpoint: url::Url,
36    http: reqwest::Client,
37}
38
39impl std::fmt::Debug for HttpClient {
40    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
41        f.debug_struct("HttpClient")
42            .field("endpoint", &self.endpoint.as_str())
43            .finish_non_exhaustive()
44    }
45}
46
47impl HttpClient {
48    /// Create a client from a base URL and target the standard ACP endpoint.
49    ///
50    /// If the URL path is empty, `/acp` is used. Otherwise `/acp` is appended
51    /// unless the path already ends with `/acp`.
52    pub fn new(base_url: impl AsRef<str>) -> Result<Self, HttpClientError> {
53        Self::with_client(base_url, reqwest::Client::new())
54    }
55
56    /// Create a client that targets the exact endpoint URL.
57    ///
58    /// Use this when connecting to a server configured with a custom
59    /// `ServerOptions::path`.
60    pub fn with_endpoint(endpoint: impl AsRef<str>) -> Result<Self, HttpClientError> {
61        Self::with_endpoint_and_client(endpoint, reqwest::Client::new())
62    }
63
64    /// Create a client with a custom HTTP client and the standard ACP endpoint.
65    ///
66    /// If the URL path is empty, `/acp` is used. Otherwise `/acp` is appended
67    /// unless the path already ends with `/acp`.
68    pub fn with_client(
69        base_url: impl AsRef<str>,
70        http: reqwest::Client,
71    ) -> Result<Self, HttpClientError> {
72        let mut endpoint = url::Url::parse(base_url.as_ref())?;
73        let path = endpoint.path().trim_end_matches('/').to_string();
74        let path = if path.is_empty() {
75            "/acp".to_string()
76        } else if path.ends_with("/acp") {
77            path
78        } else {
79            format!("{path}/acp")
80        };
81        endpoint.set_path(&path);
82        Ok(Self { endpoint, http })
83    }
84
85    /// Create a client with a custom HTTP client and exact endpoint URL.
86    ///
87    /// Use this when connecting to a server configured with a custom
88    /// `ServerOptions::path`.
89    pub fn with_endpoint_and_client(
90        endpoint: impl AsRef<str>,
91        http: reqwest::Client,
92    ) -> Result<Self, HttpClientError> {
93        let endpoint = url::Url::parse(endpoint.as_ref())?;
94        Ok(Self { endpoint, http })
95    }
96
97    fn is_websocket(&self) -> bool {
98        matches!(self.endpoint.scheme(), "ws" | "wss")
99    }
100}
101
102impl ConnectTo<Client> for HttpClient {
103    async fn connect_to(self, client: impl ConnectTo<Agent>) -> Result<(), AcpError> {
104        let (channel, transport) = ConnectTo::<Client>::into_channel_and_future(self);
105        let shutdown_tx = channel.tx.clone();
106        match futures::future::select(
107            std::pin::pin!(client.connect_to(channel)),
108            std::pin::pin!(transport),
109        )
110        .await
111        {
112            futures::future::Either::Left((result, transport)) => {
113                result?;
114
115                // Reject sends from escaped client handles while preserving
116                // messages already accepted into the channel, then let the
117                // physical transport finish those messages.
118                shutdown_tx.close_channel();
119                transport.await
120            }
121            futures::future::Either::Right((result, _)) => result,
122        }
123    }
124
125    fn into_channel_and_future(self) -> (Channel, BoxFuture<'static, Result<(), AcpError>>) {
126        let (caller, transport) = Channel::duplex();
127        (caller, Box::pin(run(self, transport)))
128    }
129}
130
131async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> {
132    if client.is_websocket() {
133        return run_ws(client, channel).await;
134    }
135    let HttpClient { endpoint, http } = client;
136    let Channel {
137        rx: mut outgoing,
138        tx: incoming,
139    } = channel;
140    let (sse_event_tx, mut sse_event_rx) = mpsc::unbounded::<SseMessage>();
141    let connection = HttpConnection::new(endpoint, http);
142    let mut state = ClientState {
143        connection: connection.clone(),
144        open_session_streams: HashSet::new(),
145        pending_requests: HashMap::new(),
146        incoming,
147    };
148    let mut lifecycle = HttpTransportLifecycle::new(connection);
149    let mut ordered_posts = PostQueue::default();
150    let mut response_posts = PostQueue::default();
151    let mut outgoing_closed = false;
152
153    let result = loop {
154        if outgoing_closed && ordered_posts.is_empty() && response_posts.is_empty() {
155            break Ok(());
156        }
157
158        let event = {
159            let outgoing_next = async {
160                if outgoing_closed {
161                    futures::future::pending().await
162                } else {
163                    outgoing.next().await
164                }
165            }
166            .fuse();
167            let sse_event_next = sse_event_rx.next().fuse();
168            let sse_failure_next = lifecycle.next_sse_failure().fuse();
169            let ordered_post_next = ordered_posts.next_completion().fuse();
170            let response_post_next = response_posts.next_completion().fuse();
171            pin_mut!(
172                outgoing_next,
173                sse_event_next,
174                sse_failure_next,
175                ordered_post_next,
176                response_post_next
177            );
178
179            futures::select! {
180                msg = outgoing_next => HttpLoopEvent::Outgoing(msg),
181                event = sse_event_next => HttpLoopEvent::SseEvent(event),
182                failure = sse_failure_next => HttpLoopEvent::SseFailure(failure),
183                post = ordered_post_next => HttpLoopEvent::Post(post),
184                post = response_post_next => HttpLoopEvent::Post(post),
185            }
186        };
187
188        let msg = match event {
189            HttpLoopEvent::Outgoing(msg) => match msg {
190                Some(Ok(msg)) => msg,
191                Some(Err(e)) => {
192                    error!("upstream channel produced error: {e}");
193                    break Err(e);
194                }
195                None => {
196                    outgoing_closed = true;
197                    continue;
198                }
199            },
200            HttpLoopEvent::SseEvent(event) => {
201                let Some(event) = event else {
202                    continue;
203                };
204                let open_session_id = state.session_to_open_for_response(&event.message);
205                state.deliver(event.message);
206                if let Some(session_id) = open_session_id {
207                    lifecycle.start_sse(Some(session_id), sse_event_tx.clone());
208                }
209                continue;
210            }
211            HttpLoopEvent::SseFailure(failure) => {
212                let scope = failure.session_id.as_deref().unwrap_or("connection");
213                error!(session_id = ?failure.session_id, error = %failure.error, "SSE stream ended");
214                break Err(AcpError::internal_error()
215                    .data(format!("{scope} SSE stream ended: {}", failure.error)));
216            }
217            HttpLoopEvent::Post(completed) => {
218                let CompletedPost {
219                    pending_request,
220                    result,
221                } = completed;
222                if let Err(e) = result {
223                    state.remove_pending_request(pending_request.as_ref());
224                    error!("POST failed: {e}");
225                    break Err(AcpError::internal_error().data(format!("POST: {e}")));
226                }
227                continue;
228            }
229        };
230
231        if state.connection.connection_id().is_none() {
232            if !is_initialize_request(&msg) {
233                break Err(AcpError::invalid_request()
234                    .data("ACP HTTP transport: first message must be `initialize`"));
235            }
236            match state.initialize(msg).await {
237                Ok(InitializeOutcome::Connected) => {
238                    lifecycle.start_sse(None, sse_event_tx.clone());
239                }
240                Ok(InitializeOutcome::Rejected) => {}
241                Err(e) => {
242                    error!("initialize failed: {e}");
243                    break Err(AcpError::internal_error().data(format!("initialize: {e}")));
244                }
245            }
246            continue;
247        }
248
249        if let Some(session_id) = session_id_from_message(&msg)
250            && state.open_session_streams.insert(session_id.clone())
251        {
252            lifecycle.start_sse(Some(session_id), sse_event_tx.clone());
253        }
254
255        let is_response = matches!(msg, RawJsonRpcMessage::Response(_));
256        match state.prepare_post(msg) {
257            // Responses answer SSE-delivered callbacks and must not be blocked
258            // behind a POST that may be waiting for that callback response.
259            Ok(post) if is_response => response_posts.push(post),
260            Ok(post) => ordered_posts.push(post),
261            Err(e) => {
262                error!("POST failed: {e}");
263                break Err(AcpError::internal_error().data(format!("POST: {e}")));
264            }
265        }
266    };
267
268    lifecycle.close().await;
269    result
270}
271
272enum HttpLoopEvent {
273    Outgoing(Option<Result<RawJsonRpcMessage, AcpError>>),
274    SseEvent(Option<SseMessage>),
275    SseFailure(SseFailure),
276    Post(CompletedPost),
277}
278
279#[derive(Debug)]
280struct SseFailure {
281    session_id: Option<String>,
282    error: String,
283}
284
285#[derive(Debug)]
286struct SseMessage {
287    message: RawJsonRpcMessage,
288}
289
290#[derive(Clone, Debug)]
291struct HttpConnection {
292    endpoint: url::Url,
293    http: reqwest::Client,
294    connection_id: Arc<StdMutex<Option<String>>>,
295}
296
297impl HttpConnection {
298    fn new(endpoint: url::Url, http: reqwest::Client) -> Self {
299        Self {
300            endpoint,
301            http,
302            connection_id: Arc::new(StdMutex::new(None)),
303        }
304    }
305
306    fn post(&self) -> reqwest::RequestBuilder {
307        self.http.post(self.endpoint.clone())
308    }
309
310    fn get(&self) -> reqwest::RequestBuilder {
311        self.http.get(self.endpoint.clone())
312    }
313
314    fn set_connection_id(&self, connection_id: String) {
315        *self.connection_id.lock().expect("mutex poisoned") = Some(connection_id);
316    }
317
318    fn connection_id(&self) -> Option<String> {
319        self.connection_id.lock().expect("mutex poisoned").clone()
320    }
321
322    fn take_connection_id(&self) -> Option<String> {
323        self.connection_id.lock().expect("mutex poisoned").take()
324    }
325
326    fn clear_connection_id(&self, expected: &str) {
327        let mut connection_id = self.connection_id.lock().expect("mutex poisoned");
328        if connection_id.as_deref() == Some(expected) {
329            *connection_id = None;
330        }
331    }
332
333    async fn close(&self) {
334        let Some(connection_id) = self.connection_id() else {
335            return;
336        };
337        Self::send_close(
338            self.http.clone(),
339            self.endpoint.clone(),
340            connection_id.clone(),
341        )
342        .await;
343        self.clear_connection_id(&connection_id);
344    }
345
346    fn spawn_close(&self) {
347        let Some(connection_id) = self.take_connection_id() else {
348            return;
349        };
350        let http = self.http.clone();
351        let endpoint = self.endpoint.clone();
352        match tokio::runtime::Handle::try_current() {
353            Ok(handle) => {
354                drop(handle.spawn(Self::send_close(http, endpoint, connection_id)));
355            }
356            Err(e) => {
357                debug!("failed to spawn HTTP DELETE: {e}");
358            }
359        }
360    }
361
362    async fn send_close(http: reqwest::Client, endpoint: url::Url, connection_id: String) {
363        if let Err(e) = http
364            .delete(endpoint)
365            .header(HEADER_CONNECTION_ID, connection_id)
366            .send()
367            .await
368        {
369            debug!("DELETE failed (ignored): {e}");
370        }
371    }
372}
373
374#[derive(Debug)]
375struct HttpTransportLifecycle {
376    connection: HttpConnection,
377    sse_tasks: SseTasks,
378}
379
380impl HttpTransportLifecycle {
381    fn new(connection: HttpConnection) -> Self {
382        Self {
383            connection,
384            sse_tasks: SseTasks::default(),
385        }
386    }
387
388    fn start_sse(&mut self, session_id: Option<String>, event_tx: UnboundedSender<SseMessage>) {
389        self.sse_tasks
390            .push(run_sse(self.connection.clone(), session_id, event_tx));
391    }
392
393    async fn next_sse_failure(&mut self) -> SseFailure {
394        self.sse_tasks.next_failure().await
395    }
396
397    async fn close(&mut self) {
398        self.connection.close().await;
399        self.sse_tasks.abort_all();
400    }
401}
402
403impl Drop for HttpTransportLifecycle {
404    fn drop(&mut self) {
405        self.sse_tasks.abort_all();
406        self.connection.spawn_close();
407    }
408}
409
410fn run_sse(
411    connection: HttpConnection,
412    session_id: Option<String>,
413    event_tx: UnboundedSender<SseMessage>,
414) -> BoxFuture<'static, SseFailure> {
415    Box::pin(async move {
416        let label = session_id.clone();
417        let error = match read_sse(connection, session_id, event_tx).await {
418            Ok(()) => "SSE stream closed".to_string(),
419            Err(e) => e,
420        };
421        warn!(session_id = ?label, "SSE stream ended: {error}");
422        SseFailure {
423            session_id: label,
424            error,
425        }
426    })
427}
428
429#[derive(Debug, Default)]
430struct SseTasks {
431    handles: FuturesUnordered<BoxFuture<'static, SseFailure>>,
432}
433
434impl SseTasks {
435    fn push(&mut self, task: BoxFuture<'static, SseFailure>) {
436        self.handles.push(task);
437    }
438
439    async fn next_failure(&mut self) -> SseFailure {
440        loop {
441            if let Some(failure) = self.handles.next().await {
442                return failure;
443            }
444            futures::future::pending::<()>().await;
445        }
446    }
447
448    fn abort_all(&mut self) {
449        self.handles = FuturesUnordered::new();
450    }
451}
452
453struct ClientState {
454    connection: HttpConnection,
455    open_session_streams: HashSet<String>,
456    pending_requests: HashMap<RequestId, String>,
457    incoming: futures::channel::mpsc::UnboundedSender<Result<RawJsonRpcMessage, AcpError>>,
458}
459
460struct PendingPost {
461    pending_request: Option<(RequestId, String)>,
462    response: BoxFuture<'static, Result<(), String>>,
463}
464
465impl PendingPost {
466    fn into_completion(self) -> BoxFuture<'static, CompletedPost> {
467        let Self {
468            pending_request,
469            response,
470        } = self;
471        async move {
472            CompletedPost {
473                pending_request,
474                result: response.await,
475            }
476        }
477        .boxed()
478    }
479}
480
481#[derive(Debug)]
482struct CompletedPost {
483    pending_request: Option<(RequestId, String)>,
484    result: Result<(), String>,
485}
486
487#[derive(Default)]
488struct PostQueue {
489    queued: VecDeque<PendingPost>,
490    in_flight: Option<BoxFuture<'static, CompletedPost>>,
491}
492
493impl PostQueue {
494    fn push(&mut self, post: PendingPost) {
495        self.queued.push_back(post);
496        self.start_next();
497    }
498
499    async fn next_completion(&mut self) -> CompletedPost {
500        loop {
501            self.start_next();
502            if let Some(in_flight) = self.in_flight.as_mut() {
503                let completed = in_flight.await;
504                self.in_flight = None;
505                return completed;
506            }
507            futures::future::pending::<()>().await;
508        }
509    }
510
511    fn start_next(&mut self) {
512        if self.in_flight.is_none()
513            && let Some(post) = self.queued.pop_front()
514        {
515            self.in_flight = Some(post.into_completion());
516        }
517    }
518
519    fn is_empty(&self) -> bool {
520        self.queued.is_empty() && self.in_flight.is_none()
521    }
522}
523
524#[derive(Clone, Copy, Debug, Eq, PartialEq)]
525enum InitializeOutcome {
526    Connected,
527    Rejected,
528}
529
530impl ClientState {
531    async fn initialize(&self, msg: RawJsonRpcMessage) -> Result<InitializeOutcome, String> {
532        let response = self
533            .connection
534            .post()
535            .header("Content-Type", "application/json")
536            .header("Accept", "application/json")
537            .json(&msg)
538            .send()
539            .await
540            .map_err(|e| e.to_string())?;
541
542        let connection_id = response
543            .headers()
544            .get(HEADER_CONNECTION_ID)
545            .and_then(|v| v.to_str().ok())
546            .map(String::from);
547        if let Some(connection_id) = &connection_id {
548            self.connection.set_connection_id(connection_id.clone());
549        }
550
551        if !response.status().is_success() {
552            let status = response.status();
553            let body = response.text().await.unwrap_or_default();
554            return Err(format!("HTTP {status}: {body}"));
555        }
556
557        let message = response
558            .json::<RawJsonRpcMessage>()
559            .await
560            .map_err(|e| e.to_string())?;
561
562        if matches!(
563            message,
564            RawJsonRpcMessage::Response(RpcResponse::Error { .. })
565        ) {
566            self.deliver(message);
567            self.connection.close().await;
568            return Ok(InitializeOutcome::Rejected);
569        }
570
571        connection_id
572            .ok_or_else(|| format!("server did not return {HEADER_CONNECTION_ID} header"))?;
573        self.deliver(message);
574        Ok(InitializeOutcome::Connected)
575    }
576
577    fn prepare_post(&mut self, msg: RawJsonRpcMessage) -> Result<PendingPost, String> {
578        let session_id = match method_for_message(&msg) {
579            Some(method) => {
580                let session_id = session_id_from_message(&msg);
581                if method_requires_session_header(method) && session_id.is_none() {
582                    return Err(format!("method `{method}` requires sessionId in params"));
583                }
584                session_id
585            }
586            None => None,
587        };
588        let connection_id = self
589            .connection
590            .connection_id()
591            .ok_or_else(|| "POST attempted before initialize".to_string())?;
592        let mut request = self
593            .connection
594            .post()
595            .header("Accept", "application/json")
596            .header(HEADER_CONNECTION_ID, connection_id)
597            .json(&msg);
598        if let Some(session_id) = session_id {
599            request = request.header(HEADER_SESSION_ID, session_id);
600        }
601
602        let pending_request = pending_request_for_message(&msg);
603        if let Some((id, method)) = &pending_request {
604            self.pending_requests.insert(id.clone(), method.clone());
605        }
606
607        let response = async move {
608            let response = request.send().await.map_err(|e| e.to_string())?;
609            if response.status().as_u16() != 202 && !response.status().is_success() {
610                let status = response.status();
611                let body = response.text().await.unwrap_or_default();
612                return Err(format!("HTTP {status}: {body}"));
613            }
614            Ok(())
615        };
616        Ok(PendingPost {
617            pending_request,
618            response: response.boxed(),
619        })
620    }
621
622    fn remove_pending_request(&mut self, pending_request: Option<&(RequestId, String)>) {
623        if let Some((id, _)) = pending_request {
624            self.pending_requests.remove(id);
625        }
626    }
627
628    fn session_to_open_for_response(&mut self, msg: &RawJsonRpcMessage) -> Option<String> {
629        let RawJsonRpcMessage::Response(response) = msg else {
630            return None;
631        };
632        let id = msg.response_id().and_then(pending_request_key)?;
633        let method = self.pending_requests.remove(&id);
634
635        if !method.as_deref().is_some_and(is_session_opening_method) {
636            return None;
637        }
638        let RpcResponse::Result { result, .. } = response else {
639            return None;
640        };
641        let session_id = result
642            .get("sessionId")
643            .and_then(|v| v.as_str())
644            .map(String::from)?;
645
646        if self.open_session_streams.insert(session_id.clone()) {
647            Some(session_id)
648        } else {
649            None
650        }
651    }
652
653    fn deliver(&self, msg: RawJsonRpcMessage) {
654        if self.incoming.unbounded_send(Ok(msg)).is_err() {
655            debug!("upstream channel closed; dropping inbound message");
656        }
657    }
658}
659
660fn is_session_opening_method(method: &str) -> bool {
661    matches!(method, "session/new" | "session/fork")
662}
663
664async fn read_sse(
665    connection: HttpConnection,
666    session_id: Option<String>,
667    event_tx: UnboundedSender<SseMessage>,
668) -> Result<(), String> {
669    let connection_id = connection
670        .connection_id()
671        .ok_or_else(|| "SSE attempted before initialize".to_string())?;
672    let mut request = connection
673        .get()
674        .header("Accept", "text/event-stream")
675        .header(HEADER_CONNECTION_ID, connection_id);
676    if let Some(session_id) = &session_id {
677        request = request.header(HEADER_SESSION_ID, session_id);
678    }
679
680    let response = request.send().await.map_err(|e| e.to_string())?;
681    if !response.status().is_success() {
682        return Err(format!("HTTP {}", response.status()));
683    }
684    trace!(session_id = ?session_id, "SSE stream open");
685
686    let mut events = eventsource_stream::EventStream::new(response.bytes_stream());
687    while let Some(event) = events.next().await {
688        let event = event.map_err(|e| e.to_string())?;
689        let payload = event.data;
690        if payload.is_empty() {
691            continue;
692        }
693        let msg = serde_json::from_str::<RawJsonRpcMessage>(&payload)
694            .map_err(|e| format!("malformed JSON-RPC payload: {e}"))?;
695
696        if event_tx
697            .unbounded_send(SseMessage { message: msg })
698            .is_err()
699        {
700            return Err("upstream channel closed".to_string());
701        }
702    }
703    Ok(())
704}
705
706fn pending_request_for_message(msg: &RawJsonRpcMessage) -> Option<(RequestId, String)> {
707    let RawJsonRpcMessage::Request(request) = msg else {
708        return None;
709    };
710    pending_request_key(&request.id).map(|id| (id, request.method.to_string()))
711}
712
713fn pending_request_key(id: &RequestId) -> Option<RequestId> {
714    match id {
715        RequestId::Null => None,
716        RequestId::Number(_) | RequestId::Str(_) => Some(id.clone()),
717    }
718}
719
720async fn run_ws(client: HttpClient, channel: Channel) -> Result<(), AcpError> {
721    let HttpClient { endpoint, .. } = client;
722
723    let (ws_stream, response) = async_tungstenite::tokio::connect_async(endpoint.as_str())
724        .await
725        .map_err(|e| AcpError::internal_error().data(format!("WebSocket connect failed: {e}")))?;
726    trace!(
727        status = %response.status(),
728        "WebSocket connection established"
729    );
730    let (ws_tx, ws_rx) = ws_stream.split();
731
732    drive_ws(ws_tx, ws_rx, channel).await
733}
734
735trait WsSink {
736    fn send(
737        &mut self,
738        message: WsMessage,
739    ) -> impl std::future::Future<Output = Result<(), String>> + Send;
740}
741
742impl<S> WsSink for async_tungstenite::WebSocketSender<S>
743where
744    S: futures::AsyncRead + futures::AsyncWrite + Unpin + Send,
745{
746    async fn send(&mut self, message: WsMessage) -> Result<(), String> {
747        async_tungstenite::WebSocketSender::send(self, message)
748            .await
749            .map_err(|error| error.to_string())
750    }
751}
752
753async fn drive_ws<Tx, Rx, RxError>(
754    mut ws_tx: Tx,
755    mut ws_rx: Rx,
756    channel: Channel,
757) -> Result<(), AcpError>
758where
759    Tx: WsSink,
760    Rx: Stream<Item = Result<WsMessage, RxError>> + Unpin,
761    RxError: std::fmt::Display,
762{
763    let Channel {
764        rx: mut outgoing,
765        tx: incoming,
766    } = channel;
767    let writer = async move {
768        while let Some(msg) = outgoing.next().await {
769            match msg {
770                Ok(msg) => {
771                    let text = match serde_json::to_string(&msg) {
772                        Ok(t) => t,
773                        Err(e) => {
774                            error!("failed to serialize outbound message: {e}");
775                            return Err(AcpError::internal_error().data(format!("serialize: {e}")));
776                        }
777                    };
778                    if let Err(e) = ws_tx.send(WsMessage::Text(text.into())).await {
779                        error!("WebSocket send failed: {e}");
780                        return Err(AcpError::internal_error().data(format!("ws send: {e}")));
781                    }
782                }
783                Err(e) => {
784                    error!("upstream channel produced error: {e}");
785                    return Err(e);
786                }
787            }
788        }
789
790        drop(ws_tx.send(WsMessage::Close(None)).await);
791        Ok(())
792    };
793
794    let reader = async move {
795        let mut discard_incoming = false;
796        loop {
797            match ws_rx.next().await {
798                Some(Ok(WsMessage::Text(text))) => {
799                    if discard_incoming {
800                        continue;
801                    }
802                    match serde_json::from_str::<RawJsonRpcMessage>(text.as_str()) {
803                        Ok(parsed) => {
804                            if incoming.unbounded_send(Ok(parsed)).is_err() {
805                                debug!(
806                                    "upstream channel closed; discarding WS input while draining output"
807                                );
808                                discard_incoming = true;
809                            }
810                        }
811                        Err(e) => {
812                            let message = format!("malformed JSON-RPC payload: {e}");
813                            warn!("WS: {message}");
814                            if incoming
815                                .unbounded_send(Err(AcpError::parse_error().data(message)))
816                                .is_err()
817                            {
818                                debug!(
819                                    "upstream channel closed; discarding WS input while draining output"
820                                );
821                                discard_incoming = true;
822                            }
823                        }
824                    }
825                }
826                Some(Ok(WsMessage::Binary(_))) => {
827                    warn!("ignoring binary WebSocket frame (ACP uses text)");
828                }
829                Some(Ok(WsMessage::Ping(_) | WsMessage::Pong(_) | WsMessage::Frame(_))) => {}
830                Some(Ok(WsMessage::Close(frame))) => {
831                    debug!("server closed WebSocket: {frame:?}");
832                    return Err(AcpError::internal_error()
833                        .data(format!("WebSocket closed by peer: {frame:?}")));
834                }
835                Some(Err(e)) => {
836                    error!("WebSocket receive error: {e}");
837                    return Err(AcpError::internal_error().data(format!("ws recv: {e}")));
838                }
839                None => {
840                    return Err(AcpError::internal_error().data("WebSocket stream ended"));
841                }
842            }
843        }
844    };
845
846    pin_mut!(writer, reader);
847    match futures::future::select(writer, reader).await {
848        futures::future::Either::Left((result, _))
849        | futures::future::Either::Right((result, _)) => result,
850    }
851}
852
853#[cfg(test)]
854mod tests {
855    use std::{
856        convert::Infallible,
857        sync::{
858            Arc,
859            atomic::{AtomicUsize, Ordering},
860        },
861        time::Duration,
862    };
863
864    use agent_client_protocol::schema::v1::RequestId;
865    use axum::{
866        Json, Router,
867        extract::{WebSocketUpgrade, ws::Message as AxumWsMessage},
868        http::{HeaderMap, HeaderValue, StatusCode},
869        response::{IntoResponse, Sse, sse::Event},
870        routing::{get, post},
871    };
872    use serde_json::json;
873    use tokio::{
874        net::TcpListener,
875        sync::Notify,
876        time::{sleep, timeout},
877    };
878
879    use super::*;
880
881    struct PostsThenExitClient {
882        finish: Arc<Notify>,
883        finished: Arc<Notify>,
884        escaped_tx: futures::channel::oneshot::Sender<
885            futures::channel::mpsc::UnboundedSender<Result<RawJsonRpcMessage, AcpError>>,
886        >,
887    }
888
889    struct QueueOutgoingThenText {
890        text: Option<WsMessage>,
891        outgoing: Option<mpsc::UnboundedSender<Result<RawJsonRpcMessage, AcpError>>>,
892    }
893
894    struct RecordingWsSink(mpsc::UnboundedSender<WsMessage>);
895
896    struct BackpressuredWsSink {
897        output: mpsc::UnboundedSender<WsMessage>,
898        started: mpsc::UnboundedSender<()>,
899        release: Option<futures::channel::oneshot::Receiver<()>>,
900    }
901
902    struct ReleaseBackpressureOnPoll {
903        started: mpsc::UnboundedReceiver<()>,
904        release: Option<futures::channel::oneshot::Sender<()>>,
905    }
906
907    impl WsSink for RecordingWsSink {
908        async fn send(&mut self, message: WsMessage) -> Result<(), String> {
909            self.0
910                .unbounded_send(message)
911                .map_err(|error| error.to_string())
912        }
913    }
914
915    impl WsSink for BackpressuredWsSink {
916        async fn send(&mut self, message: WsMessage) -> Result<(), String> {
917            self.output
918                .unbounded_send(message)
919                .map_err(|error| error.to_string())?;
920            if let Some(release) = self.release.take() {
921                self.started
922                    .unbounded_send(())
923                    .map_err(|error| error.to_string())?;
924                release
925                    .await
926                    .map_err(|_| "mock WebSocket reader did not release send".to_string())?;
927            }
928            Ok(())
929        }
930    }
931
932    impl Stream for QueueOutgoingThenText {
933        type Item = Result<WsMessage, std::io::Error>;
934
935        fn poll_next(
936            mut self: std::pin::Pin<&mut Self>,
937            _cx: &mut std::task::Context<'_>,
938        ) -> std::task::Poll<Option<Self::Item>> {
939            // Make input ready immediately after queueing output. If the output
940            // branch was polled first it was still empty, so either poll order
941            // deterministically selects this input frame first.
942            if let Some(outgoing) = self.outgoing.take() {
943                for method in ["custom/first", "custom/second"] {
944                    outgoing
945                        .unbounded_send(Ok(RawJsonRpcMessage::notification(
946                            method.to_string(),
947                            json!({}),
948                        )
949                        .unwrap()))
950                        .unwrap();
951                }
952            }
953            if let Some(text) = self.text.take() {
954                return std::task::Poll::Ready(Some(Ok(text)));
955            }
956            std::task::Poll::Pending
957        }
958    }
959
960    impl Stream for ReleaseBackpressureOnPoll {
961        type Item = Result<WsMessage, std::io::Error>;
962
963        fn poll_next(
964            mut self: std::pin::Pin<&mut Self>,
965            cx: &mut std::task::Context<'_>,
966        ) -> std::task::Poll<Option<Self::Item>> {
967            if let std::task::Poll::Ready(Some(())) =
968                std::pin::Pin::new(&mut self.started).poll_next(cx)
969                && let Some(release) = self.release.take()
970            {
971                let _result = release.send(());
972            }
973            std::task::Poll::Pending
974        }
975    }
976
977    impl ConnectTo<Agent> for PostsThenExitClient {
978        async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), AcpError> {
979            let Self {
980                finish,
981                finished,
982                escaped_tx,
983            } = self;
984            let (mut channel, transport) = agent.into_channel_and_future();
985            let client = async move {
986                escaped_tx.send(channel.tx.clone()).map_err(|_| {
987                    AcpError::internal_error().data("escaped sender observer dropped")
988                })?;
989                channel
990                    .tx
991                    .unbounded_send(Ok(RawJsonRpcMessage::request(
992                        "initialize".to_string(),
993                        json!({}),
994                        RequestId::Number(1),
995                    )
996                    .unwrap()))
997                    .map_err(|e| {
998                        AcpError::internal_error().data(format!("send initialize: {e}"))
999                    })?;
1000                channel.rx.next().await.ok_or_else(|| {
1001                    AcpError::internal_error().data("initialize response channel closed")
1002                })??;
1003
1004                for method in ["custom/first", "custom/second"] {
1005                    channel
1006                        .tx
1007                        .unbounded_send(Ok(RawJsonRpcMessage::notification(
1008                            method.to_string(),
1009                            json!({}),
1010                        )
1011                        .unwrap()))
1012                        .map_err(|e| {
1013                            AcpError::internal_error().data(format!("send {method}: {e}"))
1014                        })?;
1015                }
1016
1017                finish.notified().await;
1018                finished.notify_one();
1019                Ok(())
1020            };
1021
1022            let ((), ()) = futures::try_join!(transport, client)?;
1023            Ok(())
1024        }
1025    }
1026
1027    #[test]
1028    fn new_targets_standard_acp_endpoint() {
1029        assert_eq!(
1030            HttpClient::new("http://example.com")
1031                .unwrap()
1032                .endpoint
1033                .as_str(),
1034            "http://example.com/acp"
1035        );
1036        assert_eq!(
1037            HttpClient::new("http://example.com/proxy")
1038                .unwrap()
1039                .endpoint
1040                .as_str(),
1041            "http://example.com/proxy/acp"
1042        );
1043        assert_eq!(
1044            HttpClient::new("http://example.com/proxy/acp")
1045                .unwrap()
1046                .endpoint
1047                .as_str(),
1048            "http://example.com/proxy/acp"
1049        );
1050    }
1051
1052    #[test]
1053    fn with_endpoint_preserves_explicit_endpoint_path() {
1054        assert_eq!(
1055            HttpClient::with_endpoint("http://example.com/agent")
1056                .unwrap()
1057                .endpoint
1058                .as_str(),
1059            "http://example.com/agent"
1060        );
1061        assert_eq!(
1062            HttpClient::with_endpoint_and_client(
1063                "ws://example.com/custom/acp?token=abc",
1064                reqwest::Client::new(),
1065            )
1066            .unwrap()
1067            .endpoint
1068            .as_str(),
1069            "ws://example.com/custom/acp?token=abc"
1070        );
1071    }
1072
1073    #[tokio::test]
1074    async fn post_sends_cancel_request_without_session_header() {
1075        let (capture_tx, mut capture_rx) = tokio::sync::mpsc::unbounded_channel();
1076        let post_count = Arc::new(AtomicUsize::new(0));
1077        let app = Router::new().route(
1078            "/acp",
1079            post({
1080                let capture_tx = capture_tx.clone();
1081                let post_count = post_count.clone();
1082                move |headers: HeaderMap, Json(message): Json<RawJsonRpcMessage>| {
1083                    let capture_tx = capture_tx.clone();
1084                    let post_count = post_count.clone();
1085                    async move {
1086                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1087                            return initialize_response().await.into_response();
1088                        }
1089
1090                        capture_tx
1091                            .send((headers.get(HEADER_SESSION_ID).cloned(), message))
1092                            .unwrap();
1093                        StatusCode::ACCEPTED.into_response()
1094                    }
1095                }
1096            })
1097            .get(pending_sse)
1098            .delete(|| async { StatusCode::ACCEPTED }),
1099        );
1100        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1101        let addr = listener.local_addr().unwrap();
1102        let server = tokio::spawn(async move {
1103            axum::serve(listener, app).await.unwrap();
1104        });
1105        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1106        let (mut caller, transport) = Channel::duplex();
1107        let transport = tokio::spawn(run(client, transport));
1108
1109        caller
1110            .tx
1111            .unbounded_send(Ok(RawJsonRpcMessage::request(
1112                "initialize".to_string(),
1113                json!({}),
1114                RequestId::Number(1),
1115            )
1116            .unwrap()))
1117            .unwrap();
1118        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1119            .await
1120            .unwrap()
1121            .unwrap()
1122            .unwrap();
1123        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1124
1125        caller
1126            .tx
1127            .unbounded_send(Ok(RawJsonRpcMessage::notification(
1128                "$/cancel_request".to_string(),
1129                json!({
1130                    "requestId": 2,
1131                    "sessionId": "session-1"
1132                }),
1133            )
1134            .unwrap()))
1135            .unwrap();
1136
1137        let (session_header, message) = timeout(Duration::from_secs(1), capture_rx.recv())
1138            .await
1139            .unwrap()
1140            .unwrap();
1141        assert!(session_header.is_none());
1142        assert!(matches!(
1143            message,
1144            RawJsonRpcMessage::Notification(notification)
1145                if notification.method.as_ref() == "$/cancel_request"
1146        ));
1147
1148        drop(caller);
1149        timeout(Duration::from_secs(1), transport)
1150            .await
1151            .unwrap()
1152            .unwrap()
1153            .unwrap();
1154
1155        server.abort();
1156    }
1157
1158    #[tokio::test]
1159    async fn custom_response_with_session_id_does_not_open_session_sse() {
1160        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
1161        let response_ready = Arc::new(tokio::sync::Notify::new());
1162        let post_count = Arc::new(AtomicUsize::new(0));
1163        let app = Router::new().route(
1164            "/acp",
1165            post({
1166                let post_count = post_count.clone();
1167                let response_ready = response_ready.clone();
1168                move |Json(_message): Json<RawJsonRpcMessage>| {
1169                    let post_count = post_count.clone();
1170                    let response_ready = response_ready.clone();
1171                    async move {
1172                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1173                            return initialize_response().await.into_response();
1174                        }
1175
1176                        response_ready.notify_waiters();
1177                        StatusCode::ACCEPTED.into_response()
1178                    }
1179                }
1180            })
1181            .get({
1182                let get_tx = get_tx.clone();
1183                let response_ready = response_ready.clone();
1184                move |headers: HeaderMap| {
1185                    let get_tx = get_tx.clone();
1186                    let response_ready = response_ready.clone();
1187                    async move {
1188                        let session_header = headers
1189                            .get(HEADER_SESSION_ID)
1190                            .and_then(|value| value.to_str().ok())
1191                            .map(String::from);
1192                        get_tx.send(session_header).unwrap();
1193
1194                        let stream = async_stream::stream! {
1195                            response_ready.notified().await;
1196                            yield Ok::<_, Infallible>(sse_event(
1197                                RawJsonRpcMessage::response(
1198                                    RequestId::Number(2),
1199                                    Ok(json!({ "sessionId": "session-1" })),
1200                                ),
1201                            ));
1202                            futures::future::pending::<()>().await;
1203                        };
1204                        Sse::new(stream)
1205                    }
1206                }
1207            })
1208            .delete(|| async { StatusCode::ACCEPTED }),
1209        );
1210        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1211        let addr = listener.local_addr().unwrap();
1212        let server = tokio::spawn(async move {
1213            axum::serve(listener, app).await.unwrap();
1214        });
1215        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1216        let (mut caller, transport) = Channel::duplex();
1217        let transport = tokio::spawn(run(client, transport));
1218
1219        caller
1220            .tx
1221            .unbounded_send(Ok(RawJsonRpcMessage::request(
1222                "initialize".to_string(),
1223                json!({}),
1224                RequestId::Number(1),
1225            )
1226            .unwrap()))
1227            .unwrap();
1228        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1229            .await
1230            .unwrap()
1231            .unwrap()
1232            .unwrap();
1233        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1234
1235        let connection_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
1236            .await
1237            .unwrap()
1238            .unwrap();
1239        assert!(connection_sse_header.is_none());
1240
1241        caller
1242            .tx
1243            .unbounded_send(Ok(RawJsonRpcMessage::request(
1244                "custom/sessionish".to_string(),
1245                json!({}),
1246                RequestId::Number(2),
1247            )
1248            .unwrap()))
1249            .unwrap();
1250        let response = timeout(Duration::from_secs(1), caller.rx.next())
1251            .await
1252            .unwrap()
1253            .unwrap()
1254            .unwrap();
1255        assert!(matches!(
1256            response,
1257            RawJsonRpcMessage::Response(RpcResponse::Result {
1258                id: RequestId::Number(2),
1259                ..
1260            })
1261        ));
1262
1263        assert!(
1264            timeout(Duration::from_millis(100), get_rx.recv())
1265                .await
1266                .is_err(),
1267            "custom response must not open a session SSE stream"
1268        );
1269
1270        drop(caller);
1271        timeout(Duration::from_secs(1), transport)
1272            .await
1273            .unwrap()
1274            .unwrap()
1275            .unwrap();
1276
1277        server.abort();
1278    }
1279
1280    #[tokio::test]
1281    async fn fork_response_with_session_id_opens_session_sse() {
1282        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
1283        let response_ready = Arc::new(tokio::sync::Notify::new());
1284        let post_count = Arc::new(AtomicUsize::new(0));
1285        let app = Router::new().route(
1286            "/acp",
1287            post({
1288                let post_count = post_count.clone();
1289                let response_ready = response_ready.clone();
1290                move |Json(_message): Json<RawJsonRpcMessage>| {
1291                    let post_count = post_count.clone();
1292                    let response_ready = response_ready.clone();
1293                    async move {
1294                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
1295                            return initialize_response().await.into_response();
1296                        }
1297
1298                        response_ready.notify_waiters();
1299                        StatusCode::ACCEPTED.into_response()
1300                    }
1301                }
1302            })
1303            .get({
1304                let get_tx = get_tx.clone();
1305                let response_ready = response_ready.clone();
1306                move |headers: HeaderMap| {
1307                    let get_tx = get_tx.clone();
1308                    let response_ready = response_ready.clone();
1309                    async move {
1310                        let session_header = headers
1311                            .get(HEADER_SESSION_ID)
1312                            .and_then(|value| value.to_str().ok())
1313                            .map(String::from);
1314                        let is_connection_stream = session_header.is_none();
1315                        get_tx.send(session_header).unwrap();
1316
1317                        let stream = async_stream::stream! {
1318                            if is_connection_stream {
1319                                response_ready.notified().await;
1320                                yield Ok::<_, Infallible>(sse_event(
1321                                    RawJsonRpcMessage::response(
1322                                        RequestId::Number(2),
1323                                        Ok(json!({ "sessionId": "forked-session" })),
1324                                    ),
1325                                ));
1326                            }
1327                            futures::future::pending::<()>().await;
1328                        };
1329                        Sse::new(stream)
1330                    }
1331                }
1332            })
1333            .delete(|| async { StatusCode::ACCEPTED }),
1334        );
1335        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1336        let addr = listener.local_addr().unwrap();
1337        let server = tokio::spawn(async move {
1338            axum::serve(listener, app).await.unwrap();
1339        });
1340        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1341        let (mut caller, transport) = Channel::duplex();
1342        let transport = tokio::spawn(run(client, transport));
1343
1344        caller
1345            .tx
1346            .unbounded_send(Ok(RawJsonRpcMessage::request(
1347                "initialize".to_string(),
1348                json!({}),
1349                RequestId::Number(1),
1350            )
1351            .unwrap()))
1352            .unwrap();
1353        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1354            .await
1355            .unwrap()
1356            .unwrap()
1357            .unwrap();
1358        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1359
1360        let connection_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
1361            .await
1362            .unwrap()
1363            .unwrap();
1364        assert!(connection_sse_header.is_none());
1365
1366        caller
1367            .tx
1368            .unbounded_send(Ok(RawJsonRpcMessage::request(
1369                "session/fork".to_string(),
1370                json!({ "sessionId": "source-session" }),
1371                RequestId::Number(2),
1372            )
1373            .unwrap()))
1374            .unwrap();
1375        let source_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
1376            .await
1377            .unwrap()
1378            .unwrap();
1379        assert_eq!(source_sse_header.as_deref(), Some("source-session"));
1380
1381        let response = timeout(Duration::from_secs(1), caller.rx.next())
1382            .await
1383            .unwrap()
1384            .unwrap()
1385            .unwrap();
1386        assert!(matches!(
1387            response,
1388            RawJsonRpcMessage::Response(RpcResponse::Result {
1389                id: RequestId::Number(2),
1390                ..
1391            })
1392        ));
1393
1394        let fork_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
1395            .await
1396            .unwrap()
1397            .unwrap();
1398        assert_eq!(fork_sse_header.as_deref(), Some("forked-session"));
1399
1400        drop(caller);
1401        timeout(Duration::from_secs(1), transport)
1402            .await
1403            .unwrap()
1404            .unwrap()
1405            .unwrap();
1406
1407        server.abort();
1408    }
1409
1410    #[tokio::test]
1411    async fn client_completion_drains_ordered_posts_in_order() {
1412        let first_started = Arc::new(Notify::new());
1413        let release_first = Arc::new(Notify::new());
1414        let second_seen = Arc::new(Notify::new());
1415        let finish_client = Arc::new(Notify::new());
1416        let client_finished = Arc::new(Notify::new());
1417        let (escaped_tx, escaped_rx) = futures::channel::oneshot::channel();
1418        let app = Router::new().route(
1419            "/acp",
1420            post({
1421                let first_started = first_started.clone();
1422                let release_first = release_first.clone();
1423                let second_seen = second_seen.clone();
1424                move |Json(message): Json<RawJsonRpcMessage>| {
1425                    let first_started = first_started.clone();
1426                    let release_first = release_first.clone();
1427                    let second_seen = second_seen.clone();
1428                    async move {
1429                        if is_initialize_request(&message) {
1430                            return initialize_response().await.into_response();
1431                        }
1432
1433                        match method_for_message(&message) {
1434                            Some("custom/first") => {
1435                                first_started.notify_one();
1436                                release_first.notified().await;
1437                            }
1438                            Some("custom/second") => {
1439                                second_seen.notify_one();
1440                            }
1441                            _ => {}
1442                        }
1443                        StatusCode::ACCEPTED.into_response()
1444                    }
1445                }
1446            })
1447            .get(pending_sse)
1448            .delete(|| async { StatusCode::ACCEPTED }),
1449        );
1450        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1451        let addr = listener.local_addr().unwrap();
1452        let server = tokio::spawn(async move {
1453            axum::serve(listener, app).await.unwrap();
1454        });
1455        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1456        let mut connection = tokio::spawn(client.connect_to(PostsThenExitClient {
1457            finish: finish_client.clone(),
1458            finished: client_finished.clone(),
1459            escaped_tx,
1460        }));
1461        let escaped = timeout(Duration::from_secs(1), escaped_rx)
1462            .await
1463            .unwrap()
1464            .unwrap();
1465
1466        timeout(Duration::from_secs(1), first_started.notified())
1467            .await
1468            .unwrap();
1469        assert!(
1470            timeout(Duration::from_millis(100), second_seen.notified())
1471                .await
1472                .is_err(),
1473            "second POST must not be sent while the first POST is pending"
1474        );
1475
1476        finish_client.notify_one();
1477        timeout(Duration::from_secs(1), client_finished.notified())
1478            .await
1479            .unwrap();
1480        assert!(
1481            timeout(Duration::from_millis(100), &mut connection)
1482                .await
1483                .is_err(),
1484            "HTTP transport returned before its accepted POSTs completed"
1485        );
1486        assert!(
1487            escaped
1488                .unbounded_send(Ok(RawJsonRpcMessage::notification(
1489                    "custom/too-late".to_string(),
1490                    json!({}),
1491                )
1492                .unwrap()))
1493                .is_err(),
1494            "escaped client sender remained open after client completion"
1495        );
1496
1497        release_first.notify_one();
1498        timeout(Duration::from_secs(1), second_seen.notified())
1499            .await
1500            .unwrap();
1501
1502        timeout(Duration::from_secs(1), connection)
1503            .await
1504            .unwrap()
1505            .unwrap()
1506            .unwrap();
1507
1508        server.abort();
1509    }
1510
1511    #[tokio::test]
1512    async fn sse_continues_while_post_is_pending() {
1513        let post_started = Arc::new(Notify::new());
1514        let callback_response_seen = Arc::new(Notify::new());
1515        let sse_started = Arc::new(Notify::new());
1516        let (callback_tx, mut callback_rx) = tokio::sync::mpsc::unbounded_channel();
1517        let app = Router::new().route(
1518            "/acp",
1519            post({
1520                let post_started = post_started.clone();
1521                let callback_response_seen = callback_response_seen.clone();
1522                let callback_tx = callback_tx.clone();
1523                move |Json(message): Json<RawJsonRpcMessage>| {
1524                    let post_started = post_started.clone();
1525                    let callback_response_seen = callback_response_seen.clone();
1526                    let callback_tx = callback_tx.clone();
1527                    async move {
1528                        if is_initialize_request(&message) {
1529                            return initialize_response().await.into_response();
1530                        }
1531
1532                        match &message {
1533                            RawJsonRpcMessage::Request(request)
1534                                if request.method.as_ref() == "custom/slow" =>
1535                            {
1536                                post_started.notify_waiters();
1537                                callback_response_seen.notified().await;
1538                                StatusCode::ACCEPTED.into_response()
1539                            }
1540                            RawJsonRpcMessage::Response(
1541                                RpcResponse::Result {
1542                                    id: RequestId::Number(99),
1543                                    ..
1544                                }
1545                                | RpcResponse::Error {
1546                                    id: RequestId::Number(99),
1547                                    ..
1548                                },
1549                            ) => {
1550                                callback_tx.send(message).unwrap();
1551                                callback_response_seen.notify_waiters();
1552                                StatusCode::ACCEPTED.into_response()
1553                            }
1554                            _ => StatusCode::ACCEPTED.into_response(),
1555                        }
1556                    }
1557                }
1558            })
1559            .get({
1560                let post_started = post_started.clone();
1561                let sse_started = sse_started.clone();
1562                move || {
1563                    let post_started = post_started.clone();
1564                    let sse_started = sse_started.clone();
1565                    async move {
1566                        let stream = async_stream::stream! {
1567                            sse_started.notify_waiters();
1568                            post_started.notified().await;
1569                            yield Ok::<_, Infallible>(sse_event(
1570                                RawJsonRpcMessage::request(
1571                                    "client/callback".to_string(),
1572                                    json!({}),
1573                                    RequestId::Number(99),
1574                                )
1575                                .unwrap(),
1576                            ));
1577                            futures::future::pending::<()>().await;
1578                        };
1579                        Sse::new(stream)
1580                    }
1581                }
1582            })
1583            .delete(|| async { StatusCode::ACCEPTED }),
1584        );
1585        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1586        let addr = listener.local_addr().unwrap();
1587        let server = tokio::spawn(async move {
1588            axum::serve(listener, app).await.unwrap();
1589        });
1590        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1591        let (mut caller, transport) = Channel::duplex();
1592        let transport = tokio::spawn(run(client, transport));
1593
1594        caller
1595            .tx
1596            .unbounded_send(Ok(RawJsonRpcMessage::request(
1597                "initialize".to_string(),
1598                json!({}),
1599                RequestId::Number(1),
1600            )
1601            .unwrap()))
1602            .unwrap();
1603        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1604            .await
1605            .unwrap()
1606            .unwrap()
1607            .unwrap();
1608        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1609        timeout(Duration::from_secs(1), sse_started.notified())
1610            .await
1611            .unwrap();
1612
1613        caller
1614            .tx
1615            .unbounded_send(Ok(RawJsonRpcMessage::request(
1616                "custom/slow".to_string(),
1617                json!({}),
1618                RequestId::Number(2),
1619            )
1620            .unwrap()))
1621            .unwrap();
1622
1623        let callback = timeout(Duration::from_secs(1), caller.rx.next())
1624            .await
1625            .unwrap()
1626            .unwrap()
1627            .unwrap();
1628        assert!(matches!(
1629            callback,
1630            RawJsonRpcMessage::Request(request)
1631                if request.method.as_ref() == "client/callback"
1632                    && request.id == RequestId::Number(99)
1633        ));
1634
1635        caller
1636            .tx
1637            .unbounded_send(Ok(RawJsonRpcMessage::response(
1638                RequestId::Number(99),
1639                Ok(json!({})),
1640            )))
1641            .unwrap();
1642        let callback_response = timeout(Duration::from_secs(1), callback_rx.recv())
1643            .await
1644            .unwrap()
1645            .unwrap();
1646        assert!(matches!(
1647            callback_response,
1648            RawJsonRpcMessage::Response(RpcResponse::Result {
1649                id: RequestId::Number(99),
1650                ..
1651            })
1652        ));
1653
1654        drop(caller);
1655        timeout(Duration::from_secs(1), transport)
1656            .await
1657            .unwrap()
1658            .unwrap()
1659            .unwrap();
1660
1661        server.abort();
1662    }
1663
1664    #[tokio::test]
1665    async fn post_error_deletes_initialized_connection() {
1666        let delete_count = Arc::new(AtomicUsize::new(0));
1667        let delete_count_for_handler = delete_count.clone();
1668        let app = Router::new().route(
1669            "/acp",
1670            post(initialize_response).get(pending_sse).delete(move || {
1671                let delete_count = delete_count_for_handler.clone();
1672                async move {
1673                    delete_count.fetch_add(1, Ordering::SeqCst);
1674                    StatusCode::ACCEPTED
1675                }
1676            }),
1677        );
1678        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1679        let addr = listener.local_addr().unwrap();
1680        let server = tokio::spawn(async move {
1681            axum::serve(listener, app).await.unwrap();
1682        });
1683        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1684        let (mut caller, transport) = Channel::duplex();
1685        let transport = tokio::spawn(run(client, transport));
1686
1687        caller
1688            .tx
1689            .unbounded_send(Ok(RawJsonRpcMessage::request(
1690                "initialize".to_string(),
1691                json!({}),
1692                RequestId::Number(1),
1693            )
1694            .unwrap()))
1695            .unwrap();
1696        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1697            .await
1698            .unwrap()
1699            .unwrap()
1700            .unwrap();
1701        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1702
1703        caller
1704            .tx
1705            .unbounded_send(Ok(RawJsonRpcMessage::request(
1706                "session/prompt".to_string(),
1707                json!({}),
1708                RequestId::Number(2),
1709            )
1710            .unwrap()))
1711            .unwrap();
1712        let error = timeout(Duration::from_secs(1), transport)
1713            .await
1714            .unwrap()
1715            .unwrap()
1716            .unwrap_err();
1717
1718        assert!(error.to_string().contains("POST"));
1719        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
1720
1721        server.abort();
1722    }
1723
1724    #[tokio::test]
1725    async fn connection_sse_disconnect_fails_transport() {
1726        let delete_count = Arc::new(AtomicUsize::new(0));
1727        let delete_count_for_handler = delete_count.clone();
1728        let app = Router::new().route(
1729            "/acp",
1730            post(initialize_response).get(closed_sse).delete(move || {
1731                let delete_count = delete_count_for_handler.clone();
1732                async move {
1733                    delete_count.fetch_add(1, Ordering::SeqCst);
1734                    StatusCode::ACCEPTED
1735                }
1736            }),
1737        );
1738        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1739        let addr = listener.local_addr().unwrap();
1740        let server = tokio::spawn(async move {
1741            axum::serve(listener, app).await.unwrap();
1742        });
1743        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1744        let (mut caller, transport) = Channel::duplex();
1745        let transport = tokio::spawn(run(client, transport));
1746
1747        caller
1748            .tx
1749            .unbounded_send(Ok(RawJsonRpcMessage::request(
1750                "initialize".to_string(),
1751                json!({}),
1752                RequestId::Number(1),
1753            )
1754            .unwrap()))
1755            .unwrap();
1756        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1757            .await
1758            .unwrap()
1759            .unwrap()
1760            .unwrap();
1761        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1762
1763        let error = timeout(Duration::from_secs(1), transport)
1764            .await
1765            .unwrap()
1766            .unwrap()
1767            .unwrap_err();
1768
1769        assert!(error.to_string().contains("SSE"));
1770        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
1771
1772        server.abort();
1773    }
1774
1775    #[tokio::test]
1776    async fn malformed_sse_json_fails_transport() {
1777        let delete_count = Arc::new(AtomicUsize::new(0));
1778        let delete_count_for_handler = delete_count.clone();
1779        let app = Router::new().route(
1780            "/acp",
1781            post(initialize_response)
1782                .get(malformed_sse)
1783                .delete(move || {
1784                    let delete_count = delete_count_for_handler.clone();
1785                    async move {
1786                        delete_count.fetch_add(1, Ordering::SeqCst);
1787                        StatusCode::ACCEPTED
1788                    }
1789                }),
1790        );
1791        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1792        let addr = listener.local_addr().unwrap();
1793        let server = tokio::spawn(async move {
1794            axum::serve(listener, app).await.unwrap();
1795        });
1796        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1797        let (mut caller, transport) = Channel::duplex();
1798        let transport = tokio::spawn(run(client, transport));
1799
1800        caller
1801            .tx
1802            .unbounded_send(Ok(RawJsonRpcMessage::request(
1803                "initialize".to_string(),
1804                json!({}),
1805                RequestId::Number(1),
1806            )
1807            .unwrap()))
1808            .unwrap();
1809        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
1810            .await
1811            .unwrap()
1812            .unwrap()
1813            .unwrap();
1814        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
1815
1816        let error = timeout(Duration::from_secs(1), transport)
1817            .await
1818            .unwrap()
1819            .unwrap()
1820            .unwrap_err();
1821
1822        assert!(error.to_string().contains("malformed JSON-RPC payload"));
1823        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
1824
1825        server.abort();
1826    }
1827
1828    #[tokio::test]
1829    async fn malformed_ws_json_reports_parse_error_and_continues() {
1830        let app = Router::new().route("/acp", get(malformed_then_valid_ws));
1831        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1832        let addr = listener.local_addr().unwrap();
1833        let server = tokio::spawn(async move {
1834            axum::serve(listener, app).await.unwrap();
1835        });
1836        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
1837        let (mut caller, transport) = Channel::duplex();
1838        let transport = tokio::spawn(run(client, transport));
1839
1840        let error = timeout(Duration::from_secs(1), caller.rx.next())
1841            .await
1842            .unwrap()
1843            .unwrap()
1844            .unwrap_err();
1845        assert!(error.to_string().contains("malformed JSON-RPC payload"));
1846
1847        let message = timeout(Duration::from_secs(1), caller.rx.next())
1848            .await
1849            .unwrap()
1850            .unwrap()
1851            .unwrap();
1852        assert!(matches!(message, RawJsonRpcMessage::Response(_)));
1853
1854        drop(caller);
1855        timeout(Duration::from_secs(1), transport)
1856            .await
1857            .unwrap()
1858            .unwrap()
1859            .unwrap();
1860
1861        server.abort();
1862    }
1863
1864    #[tokio::test]
1865    async fn websocket_drain_discards_incoming_after_receiver_closes() {
1866        let (caller, transport) = Channel::duplex();
1867        let Channel {
1868            tx: outgoing,
1869            rx: incoming,
1870        } = caller;
1871        drop(incoming);
1872
1873        let inbound =
1874            RawJsonRpcMessage::notification("custom/inbound".to_string(), json!({})).unwrap();
1875        let inbound = WsMessage::Text(serde_json::to_string(&inbound).unwrap().into());
1876        let ws_rx = QueueOutgoingThenText {
1877            text: Some(inbound),
1878            outgoing: Some(outgoing),
1879        };
1880        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
1881        timeout(
1882            Duration::from_secs(1),
1883            drive_ws(RecordingWsSink(ws_output_tx), ws_rx, transport),
1884        )
1885        .await
1886        .unwrap()
1887        .unwrap();
1888        let mut frames = Vec::new();
1889        while let Some(frame) = ws_output.next().await {
1890            frames.push(frame);
1891        }
1892
1893        let messages = frames
1894            .iter()
1895            .filter_map(|frame| match frame {
1896                WsMessage::Text(text) => {
1897                    Some(serde_json::from_str::<RawJsonRpcMessage>(text.as_str()).unwrap())
1898                }
1899                _ => None,
1900            })
1901            .collect::<Vec<_>>();
1902        let methods = messages
1903            .iter()
1904            .filter_map(method_for_message)
1905            .collect::<Vec<_>>();
1906        assert_eq!(methods, ["custom/first", "custom/second"]);
1907        assert!(matches!(frames.last(), Some(WsMessage::Close(None))));
1908    }
1909
1910    #[tokio::test]
1911    async fn websocket_reader_runs_while_send_is_backpressured() {
1912        let (caller, transport) = Channel::duplex();
1913        let Channel {
1914            tx: outgoing,
1915            rx: incoming,
1916        } = caller;
1917        drop(incoming);
1918        outgoing
1919            .unbounded_send(Ok(RawJsonRpcMessage::notification(
1920                "custom/queued".to_string(),
1921                json!({}),
1922            )
1923            .unwrap()))
1924            .unwrap();
1925        drop(outgoing);
1926
1927        let (started_tx, started_rx) = mpsc::unbounded();
1928        let (release_tx, release_rx) = futures::channel::oneshot::channel();
1929        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
1930        let ws_tx = BackpressuredWsSink {
1931            output: ws_output_tx,
1932            started: started_tx,
1933            release: Some(release_rx),
1934        };
1935        let ws_rx = ReleaseBackpressureOnPoll {
1936            started: started_rx,
1937            release: Some(release_tx),
1938        };
1939
1940        timeout(Duration::from_secs(1), drive_ws(ws_tx, ws_rx, transport))
1941            .await
1942            .expect("WebSocket reader was not polled while its writer was backpressured")
1943            .unwrap();
1944        let frames = ws_output.by_ref().collect::<Vec<_>>().await;
1945
1946        let WsMessage::Text(text) = &frames[0] else {
1947            panic!("queued message was not sent as WebSocket text");
1948        };
1949        let message = serde_json::from_str::<RawJsonRpcMessage>(text.as_str()).unwrap();
1950        assert_eq!(method_for_message(&message), Some("custom/queued"));
1951        assert!(matches!(frames.get(1), Some(WsMessage::Close(None))));
1952        assert_eq!(frames.len(), 2);
1953    }
1954
1955    #[tokio::test]
1956    async fn peer_ws_close_fails_transport() {
1957        let app = Router::new().route("/acp", get(close_ws));
1958        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1959        let addr = listener.local_addr().unwrap();
1960        let server = tokio::spawn(async move {
1961            axum::serve(listener, app).await.unwrap();
1962        });
1963        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
1964        let (_caller, transport) = Channel::duplex();
1965        let transport = tokio::spawn(run(client, transport));
1966
1967        let error = timeout(Duration::from_secs(1), transport)
1968            .await
1969            .unwrap()
1970            .unwrap()
1971            .unwrap_err();
1972        assert!(error.to_string().contains("WebSocket closed by peer"));
1973
1974        server.abort();
1975    }
1976
1977    #[tokio::test]
1978    async fn dropped_transport_future_deletes_initialized_connection() {
1979        let delete_count = Arc::new(AtomicUsize::new(0));
1980        let delete_count_for_handler = delete_count.clone();
1981        let app = Router::new().route(
1982            "/acp",
1983            post(initialize_response).get(pending_sse).delete(move || {
1984                let delete_count = delete_count_for_handler.clone();
1985                async move {
1986                    delete_count.fetch_add(1, Ordering::SeqCst);
1987                    StatusCode::ACCEPTED
1988                }
1989            }),
1990        );
1991        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
1992        let addr = listener.local_addr().unwrap();
1993        let server = tokio::spawn(async move {
1994            axum::serve(listener, app).await.unwrap();
1995        });
1996        let client = HttpClient::new(format!("http://{addr}")).unwrap();
1997        let (mut caller, transport) = Channel::duplex();
1998        let mut transport = Box::pin(run(client, transport));
1999
2000        caller
2001            .tx
2002            .unbounded_send(Ok(RawJsonRpcMessage::request(
2003                "initialize".to_string(),
2004                json!({}),
2005                RequestId::Number(1),
2006            )
2007            .unwrap()))
2008            .unwrap();
2009        let init_response = timeout(Duration::from_secs(1), async {
2010            tokio::select! {
2011                result = &mut transport => {
2012                    panic!("transport ended before initialize response: {result:?}");
2013                }
2014                msg = caller.rx.next() => {
2015                    msg.unwrap().unwrap()
2016                }
2017            }
2018        })
2019        .await
2020        .unwrap();
2021        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2022
2023        drop(transport);
2024        wait_for_delete(&delete_count).await;
2025
2026        server.abort();
2027    }
2028
2029    #[tokio::test]
2030    async fn dropped_transport_during_close_retries_delete() {
2031        let delete_count = Arc::new(AtomicUsize::new(0));
2032        let delete_count_for_handler = delete_count.clone();
2033        let release_delete = Arc::new(Notify::new());
2034        let release_delete_for_handler = release_delete.clone();
2035        let app = Router::new().route(
2036            "/acp",
2037            post(initialize_response).get(pending_sse).delete(move || {
2038                let delete_count = delete_count_for_handler.clone();
2039                let release_delete = release_delete_for_handler.clone();
2040                async move {
2041                    delete_count.fetch_add(1, Ordering::SeqCst);
2042                    release_delete.notified().await;
2043                    StatusCode::ACCEPTED
2044                }
2045            }),
2046        );
2047        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2048        let addr = listener.local_addr().unwrap();
2049        let server = tokio::spawn(async move {
2050            axum::serve(listener, app).await.unwrap();
2051        });
2052        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2053        let (mut caller, transport) = Channel::duplex();
2054        let transport = tokio::spawn(run(client, transport));
2055
2056        caller
2057            .tx
2058            .unbounded_send(Ok(RawJsonRpcMessage::request(
2059                "initialize".to_string(),
2060                json!({}),
2061                RequestId::Number(1),
2062            )
2063            .unwrap()))
2064            .unwrap();
2065        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2066            .await
2067            .unwrap()
2068            .unwrap()
2069            .unwrap();
2070        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2071
2072        drop(caller);
2073        wait_for_delete_count(&delete_count, 1).await;
2074        transport.abort();
2075        wait_for_delete_count(&delete_count, 2).await;
2076        release_delete.notify_waiters();
2077        drop(transport.await);
2078
2079        server.abort();
2080    }
2081
2082    #[tokio::test]
2083    async fn initialize_error_without_connection_id_is_delivered_without_sse() {
2084        let get_count = Arc::new(AtomicUsize::new(0));
2085        let get_count_for_handler = get_count.clone();
2086        let app = Router::new().route(
2087            "/acp",
2088            post(initialize_error_response).get(move || {
2089                let get_count = get_count_for_handler.clone();
2090                async move {
2091                    get_count.fetch_add(1, Ordering::SeqCst);
2092                    pending_sse().await
2093                }
2094            }),
2095        );
2096        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2097        let addr = listener.local_addr().unwrap();
2098        let server = tokio::spawn(async move {
2099            axum::serve(listener, app).await.unwrap();
2100        });
2101        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2102        let (mut caller, transport) = Channel::duplex();
2103        let transport = tokio::spawn(run(client, transport));
2104
2105        caller
2106            .tx
2107            .unbounded_send(Ok(RawJsonRpcMessage::request(
2108                "initialize".to_string(),
2109                json!({}),
2110                RequestId::Number(1),
2111            )
2112            .unwrap()))
2113            .unwrap();
2114        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2115            .await
2116            .unwrap()
2117            .unwrap()
2118            .unwrap();
2119
2120        assert!(matches!(
2121            init_response,
2122            RawJsonRpcMessage::Response(RpcResponse::Error {
2123                id: RequestId::Number(1),
2124                ..
2125            })
2126        ));
2127        assert_eq!(get_count.load(Ordering::SeqCst), 0);
2128
2129        drop(caller);
2130        timeout(Duration::from_secs(1), transport)
2131            .await
2132            .unwrap()
2133            .unwrap()
2134            .unwrap();
2135
2136        server.abort();
2137    }
2138
2139    #[tokio::test]
2140    async fn malformed_initialize_body_with_connection_id_is_deleted() {
2141        let delete_count = Arc::new(AtomicUsize::new(0));
2142        let delete_count_for_handler = delete_count.clone();
2143        let app = Router::new().route(
2144            "/acp",
2145            post(malformed_initialize_response).delete(move || {
2146                let delete_count = delete_count_for_handler.clone();
2147                async move {
2148                    delete_count.fetch_add(1, Ordering::SeqCst);
2149                    StatusCode::ACCEPTED
2150                }
2151            }),
2152        );
2153        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2154        let addr = listener.local_addr().unwrap();
2155        let server = tokio::spawn(async move {
2156            axum::serve(listener, app).await.unwrap();
2157        });
2158        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2159        let (caller, transport) = Channel::duplex();
2160        let transport = tokio::spawn(run(client, transport));
2161
2162        caller
2163            .tx
2164            .unbounded_send(Ok(RawJsonRpcMessage::request(
2165                "initialize".to_string(),
2166                json!({}),
2167                RequestId::Number(1),
2168            )
2169            .unwrap()))
2170            .unwrap();
2171        let error = timeout(Duration::from_secs(1), transport)
2172            .await
2173            .unwrap()
2174            .unwrap()
2175            .unwrap_err();
2176
2177        assert!(error.to_string().contains("initialize"));
2178        wait_for_delete(&delete_count).await;
2179
2180        server.abort();
2181    }
2182
2183    async fn wait_for_delete(delete_count: &AtomicUsize) {
2184        wait_for_delete_count(delete_count, 1).await;
2185        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
2186    }
2187
2188    async fn wait_for_delete_count(delete_count: &AtomicUsize, expected: usize) {
2189        timeout(Duration::from_secs(1), async {
2190            loop {
2191                if delete_count.load(Ordering::SeqCst) >= expected {
2192                    break;
2193                }
2194                sleep(Duration::from_millis(10)).await;
2195            }
2196        })
2197        .await
2198        .unwrap();
2199    }
2200
2201    async fn initialize_response() -> impl IntoResponse {
2202        let mut headers = HeaderMap::new();
2203        headers.insert(HEADER_CONNECTION_ID, HeaderValue::from_static("conn-1"));
2204        (
2205            StatusCode::OK,
2206            headers,
2207            Json(RawJsonRpcMessage::response(
2208                RequestId::Number(1),
2209                Ok(json!({})),
2210            )),
2211        )
2212    }
2213
2214    async fn initialize_error_response() -> Json<RawJsonRpcMessage> {
2215        Json(RawJsonRpcMessage::response(
2216            RequestId::Number(1),
2217            Err(AcpError::invalid_request().data("initialize rejected")),
2218        ))
2219    }
2220
2221    async fn malformed_initialize_response() -> impl IntoResponse {
2222        let mut headers = HeaderMap::new();
2223        headers.insert(HEADER_CONNECTION_ID, HeaderValue::from_static("conn-1"));
2224        (StatusCode::OK, headers, "{not json")
2225    }
2226
2227    async fn pending_sse() -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
2228        Sse::new(futures::stream::pending())
2229    }
2230
2231    fn sse_event(message: RawJsonRpcMessage) -> Event {
2232        Event::default().data(serde_json::to_string(&message).unwrap())
2233    }
2234
2235    async fn malformed_sse() -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
2236        let invalid = futures::stream::once(async {
2237            Ok::<_, Infallible>(Event::default().data("{not json"))
2238        });
2239        Sse::new(invalid.chain(futures::stream::pending()))
2240    }
2241
2242    async fn malformed_then_valid_ws(ws: WebSocketUpgrade) -> impl IntoResponse {
2243        ws.on_upgrade(|mut socket| async move {
2244            drop(socket.send(AxumWsMessage::Text("{not json".into())).await);
2245            let valid = serde_json::to_string(&RawJsonRpcMessage::response(
2246                RequestId::Number(1),
2247                Ok(json!({})),
2248            ))
2249            .unwrap();
2250            drop(socket.send(AxumWsMessage::Text(valid.into())).await);
2251            futures::future::pending::<()>().await;
2252        })
2253    }
2254
2255    async fn close_ws(ws: WebSocketUpgrade) -> impl IntoResponse {
2256        ws.on_upgrade(|mut socket| async move {
2257            drop(socket.send(AxumWsMessage::Close(None)).await);
2258        })
2259    }
2260
2261    async fn closed_sse() -> StatusCode {
2262        StatusCode::OK
2263    }
2264}