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