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, ConnectionDriver, Error as AcpError, RawJsonRpcMessage,
8    RawJsonRpcResponse as RpcResponse, TransportBatchEntry, TransportFrame, schema::v1::RequestId,
9};
10use async_tungstenite::tungstenite::Message as WsMessage;
11use futures::{
12    Stream, StreamExt,
13    channel::{
14        mpsc::{self, UnboundedSender},
15        oneshot,
16    },
17    future::{BoxFuture, FutureExt},
18    pin_mut,
19    stream::FuturesUnordered,
20};
21use thiserror::Error;
22use tracing::{debug, error, trace, warn};
23
24use crate::protocol::{
25    HEADER_CONNECTION_ID, HEADER_SESSION_ID, is_initialize_request, is_response_only_shape,
26    method_for_message, method_requires_session_header, session_id_from_message,
27};
28
29#[derive(Debug, Error)]
30pub enum HttpClientError {
31    #[error("invalid URL: {0}")]
32    InvalidUrl(#[from] url::ParseError),
33    #[error("unsupported URL scheme: {0}; expected http, https, ws, or wss")]
34    UnsupportedScheme(String),
35    #[error(
36        "WebSocket URLs require HttpClient::builder or builder_with_endpoint; a prebuilt reqwest client cannot enforce WebSocket connection policies"
37    )]
38    WebSocketRequiresBuilder,
39    #[error("failed to build HTTP client: {0}")]
40    Reqwest(#[from] reqwest::Error),
41}
42
43/// An endpoint-bound ACP transport using HTTP/SSE or WebSocket.
44///
45/// Cloning shares the underlying HTTP connection pool. Each connection has
46/// independent ACP transport state.
47#[derive(Clone)]
48pub struct HttpClient {
49    endpoint: url::Url,
50    http: reqwest::Client,
51}
52
53/// Configures an [`HttpClient`] before its underlying HTTP client is built.
54///
55/// Use [`HttpClient::builder`] for a base URL or
56/// [`HttpClient::builder_with_endpoint`] for an exact endpoint.
57#[must_use = "the builder must be built to create an HTTP client"]
58pub struct HttpClientBuilder {
59    endpoint: Result<url::Url, HttpClientError>,
60    http: reqwest::ClientBuilder,
61}
62
63impl std::fmt::Debug for HttpClient {
64    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
65        f.debug_struct("HttpClient")
66            .field("endpoint", &self.endpoint.as_str())
67            .finish_non_exhaustive()
68    }
69}
70
71impl std::fmt::Debug for HttpClientBuilder {
72    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
73        f.debug_struct("HttpClientBuilder")
74            .field("endpoint", &self.endpoint)
75            .finish_non_exhaustive()
76    }
77}
78
79impl HttpClient {
80    /// Create a client from a base URL and target the standard ACP endpoint.
81    ///
82    /// If the URL path is empty, `/acp` is used. Otherwise `/acp` is appended
83    /// unless the path already ends with `/acp`.
84    ///
85    /// Use [`Self::builder`] to customize headers, TLS, proxies, or timeouts.
86    pub fn new(base_url: impl AsRef<str>) -> Result<Self, HttpClientError> {
87        Self::builder(base_url).build()
88    }
89
90    /// Create a client that targets the exact endpoint URL.
91    ///
92    /// Use this when connecting to a server configured with a custom
93    /// `ServerOptions::path`. Use [`Self::builder_with_endpoint`] to also
94    /// customize the HTTP client configuration.
95    pub fn with_endpoint(endpoint: impl AsRef<str>) -> Result<Self, HttpClientError> {
96        Self::builder_with_endpoint(endpoint).build()
97    }
98
99    /// Configure a client from a base URL targeting the standard ACP endpoint.
100    ///
101    /// Uses the same path normalization as [`Self::new`]. Invalid URLs and
102    /// unsupported schemes are reported by [`HttpClientBuilder::build`].
103    ///
104    /// ```
105    /// use std::time::Duration;
106    /// use agent_client_protocol_http::HttpClient;
107    ///
108    /// let transport = HttpClient::builder("wss://agent.example")
109    ///     .configure_http(|http| {
110    ///         http.connect_timeout(Duration::from_secs(5))
111    ///             .timeout(Duration::from_secs(10))
112    ///     })
113    ///     .build()?;
114    /// # Ok::<(), agent_client_protocol_http::HttpClientError>(())
115    /// ```
116    pub fn builder(base_url: impl AsRef<str>) -> HttpClientBuilder {
117        HttpClientBuilder {
118            endpoint: parse_base_url(base_url.as_ref()),
119            http: reqwest::Client::builder(),
120        }
121    }
122
123    /// Configure a client targeting an exact endpoint URL without changing its path.
124    ///
125    /// Use this when connecting to a server configured with a custom
126    /// `ServerOptions::path`. Invalid URLs and unsupported schemes are reported
127    /// by [`HttpClientBuilder::build`].
128    pub fn builder_with_endpoint(endpoint: impl AsRef<str>) -> HttpClientBuilder {
129        HttpClientBuilder {
130            endpoint: parse_endpoint(endpoint.as_ref()),
131            http: reqwest::Client::builder(),
132        }
133    }
134
135    /// Reuse an existing reqwest client for HTTP/SSE at an exact endpoint URL.
136    ///
137    /// The path is not changed: include `/acp` or the server's custom path.
138    /// This preserves the supplied client's configuration and connection pool.
139    ///
140    /// Only `http://` and `https://` URLs are accepted. For WebSockets, use
141    /// [`Self::builder`] or [`Self::builder_with_endpoint`] so this transport can
142    /// configure HTTP/1.1 and disable handshake redirects before building the
143    /// underlying client.
144    pub fn from_http_client(
145        endpoint: impl AsRef<str>,
146        http: reqwest::Client,
147    ) -> Result<Self, HttpClientError> {
148        let endpoint = parse_endpoint(endpoint.as_ref())?;
149        if is_websocket_url(&endpoint) {
150            return Err(HttpClientError::WebSocketRequiresBuilder);
151        }
152        Ok(Self { endpoint, http })
153    }
154
155    /// Reuse an existing HTTP/SSE client with the standard ACP endpoint.
156    ///
157    /// Preserves the same base-URL path normalization as [`Self::new`].
158    /// WebSocket URLs return [`HttpClientError::WebSocketRequiresBuilder`]
159    /// before any network I/O; migrate those calls to [`Self::builder`].
160    #[deprecated(
161        note = "Use builder(...).configure_http(...).build(), or from_http_client with an exact HTTP/SSE endpoint"
162    )]
163    pub fn with_client(
164        base_url: impl AsRef<str>,
165        http: reqwest::Client,
166    ) -> Result<Self, HttpClientError> {
167        Self::from_http_client(parse_base_url(base_url.as_ref())?, http)
168    }
169
170    /// Reuse an existing HTTP/SSE client at an exact endpoint.
171    ///
172    /// WebSocket URLs return [`HttpClientError::WebSocketRequiresBuilder`]
173    /// before any network I/O; migrate those calls to
174    /// [`Self::builder_with_endpoint`].
175    #[deprecated(
176        note = "Use builder_with_endpoint(...).configure_http(...).build(), or from_http_client for an existing HTTP/SSE client"
177    )]
178    pub fn with_endpoint_and_client(
179        endpoint: impl AsRef<str>,
180        http: reqwest::Client,
181    ) -> Result<Self, HttpClientError> {
182        Self::from_http_client(endpoint, http)
183    }
184
185    fn is_websocket(&self) -> bool {
186        is_websocket_url(&self.endpoint)
187    }
188}
189
190impl HttpClientBuilder {
191    /// Customize the HTTP client used for HTTP/SSE or the WebSocket handshake.
192    ///
193    /// Each call transforms the current configuration, retaining previous changes.
194    /// Configure default headers, proxies, DNS, trust roots, client certificates,
195    /// and timeouts through reqwest's builder rather than building a client first.
196    ///
197    /// For WebSockets, [`Self::build`] overrides the HTTP version preference with
198    /// HTTP/1.1 and disables redirects. HTTP/SSE retains the supplied settings.
199    /// Subprotocols and extensions are not negotiated, even if custom default
200    /// headers request them.
201    ///
202    /// # Proxy headers
203    ///
204    /// Do not set `Host`, `Connection`, `Upgrade`, or any `Sec-WebSocket-*` header
205    /// through [`reqwest::Proxy::headers`] for WebSocket clients. The SDK sets
206    /// `Connection`, `Upgrade`, `Sec-WebSocket-Version`, and `Sec-WebSocket-Key`
207    /// over `default_headers`, but reqwest can overwrite them
208    /// afterward with proxy headers on plain `ws://` connections. The SDK cannot
209    /// inspect or reject this opaque proxy configuration at build time.
210    /// Response validation still runs before sending ACP data; it does not prove
211    /// that proxy configuration left every request header unchanged.
212    ///
213    /// # Defaults
214    ///
215    /// Both transports use reqwest's proxy discovery and TLS verification defaults.
216    /// This changes unconfigured WebSockets from direct connections with bundled
217    /// WebPKI roots to environment proxies (and system proxies when enabled) and,
218    /// with the SDK's rustls configuration, platform certificate verification.
219    /// Use `http.no_proxy()` to connect directly. Use `tls_certs_only` for an
220    /// explicit root set, or `tls_certs_merge` to add roots to the default trust:
221    ///
222    /// ```
223    /// use agent_client_protocol_http::HttpClient;
224    ///
225    /// fn direct_with_root(root_pem: &[u8]) -> Result<HttpClient, Box<dyn std::error::Error>> {
226    ///     let root = reqwest::Certificate::from_pem(root_pem)?;
227    ///     Ok(HttpClient::builder("wss://agent.example")
228    ///         .configure_http(|http| http.no_proxy().tls_certs_only([root]))
229    ///         .build()?)
230    /// }
231    /// ```
232    ///
233    /// # Timeouts
234    ///
235    /// Reqwest request/read timeouts apply to the WebSocket opening handshake,
236    /// not the lifetime of the upgraded socket. For HTTP/SSE, they retain their
237    /// normal reqwest request/body semantics, including long-lived SSE bodies.
238    ///
239    /// # Preconfigured TLS
240    ///
241    /// Prefer reqwest's TLS options for custom roots and identities. If using
242    /// `tls_backend_preconfigured`, its ALPN configuration must itself use
243    /// HTTP/1.1 for WebSockets: reqwest cannot rewrite a preconfigured backend's
244    /// ALPN. Incompatible negotiation is rejected before transmitting ACP data.
245    pub fn configure_http(
246        mut self,
247        configure: impl FnOnce(reqwest::ClientBuilder) -> reqwest::ClientBuilder,
248    ) -> Self {
249        self.http = configure(self.http);
250        self
251    }
252
253    /// Build a client, applying the selected transport's connection policies.
254    ///
255    /// Accepts `http`, `https`, `ws`, and `wss` URLs. No connection is opened
256    /// until the resulting [`HttpClient`] is connected through [`ConnectTo`].
257    pub fn build(self) -> Result<HttpClient, HttpClientError> {
258        let endpoint = self.endpoint?;
259        let http = if is_websocket_url(&endpoint) {
260            self.http
261                .http1_only()
262                .redirect(reqwest::redirect::Policy::none())
263        } else {
264            self.http
265        }
266        .build()?;
267        Ok(HttpClient { endpoint, http })
268    }
269}
270
271fn parse_base_url(base_url: &str) -> Result<url::Url, HttpClientError> {
272    let mut endpoint = parse_endpoint(base_url)?;
273    let path = endpoint.path().trim_end_matches('/');
274    let path = if path.is_empty() {
275        "/acp".to_string()
276    } else if path.ends_with("/acp") {
277        path.to_string()
278    } else {
279        format!("{path}/acp")
280    };
281    endpoint.set_path(&path);
282    Ok(endpoint)
283}
284
285fn parse_endpoint(endpoint: &str) -> Result<url::Url, HttpClientError> {
286    let endpoint = url::Url::parse(endpoint)?;
287    match endpoint.scheme() {
288        "http" | "https" | "ws" | "wss" => Ok(endpoint),
289        scheme => Err(HttpClientError::UnsupportedScheme(scheme.to_string())),
290    }
291}
292
293fn is_websocket_url(endpoint: &url::Url) -> bool {
294    matches!(endpoint.scheme(), "ws" | "wss")
295}
296
297impl ConnectTo<Client> for HttpClient {
298    async fn connect_to(self, client: impl ConnectTo<Agent>) -> Result<(), AcpError> {
299        let (channel, transport) = ConnectTo::<Client>::into_channel_and_future(self);
300        let transport = transport.expect("HttpClient owns its physical transport driver");
301        match futures::future::select(
302            std::pin::pin!(client.connect_to(channel)),
303            std::pin::pin!(transport),
304        )
305        .await
306        {
307            futures::future::Either::Left((result, mut transport)) => {
308                result?;
309
310                // Reject sends from escaped client handles while preserving
311                // messages already accepted into the channel, then let the
312                // physical transport finish those messages.
313                assert!(transport.request_finish());
314                transport.await
315            }
316            futures::future::Either::Right((result, _)) => result,
317        }
318    }
319
320    fn into_channel_and_future(self) -> (Channel, Option<ConnectionDriver>) {
321        let (caller, transport) = Channel::duplex();
322        let (finish_tx, finish_rx) = oneshot::channel();
323        let driver = ConnectionDriver::with_finish(
324            run_with_finish(self, transport, Some(finish_rx)),
325            move || {
326                // The core has handed off its accepted output before requesting
327                // finish. Seal the producer, including escaped sender clones, and
328                // let run drain queued frames and complete the physical transport.
329                let _ = finish_tx.send(());
330            },
331        );
332        (caller, Some(driver))
333    }
334}
335
336// A finish signal seals the receiver rather than retaining a producer clone:
337// ordinary producer EOF still works, and losing the hook is not a finish request.
338fn finishable_outgoing(
339    mut outgoing: mpsc::UnboundedReceiver<TransportFrame>,
340    mut finish: Option<oneshot::Receiver<()>>,
341) -> impl Stream<Item = TransportFrame> + Unpin + Send {
342    futures::stream::poll_fn(move |cx| {
343        if let Some(signal) = finish.as_mut()
344            && let std::task::Poll::Ready(result) = signal.poll_unpin(cx)
345        {
346            if result.is_ok() {
347                outgoing.close();
348            }
349            finish = None;
350        }
351        outgoing.poll_next_unpin(cx)
352    })
353}
354
355#[cfg(test)]
356async fn run(client: HttpClient, channel: Channel) -> Result<(), AcpError> {
357    run_with_finish(client, channel, None).await
358}
359
360async fn run_with_finish(
361    client: HttpClient,
362    channel: Channel,
363    finish: Option<oneshot::Receiver<()>>,
364) -> Result<(), AcpError> {
365    if client.is_websocket() {
366        return run_ws(client, channel, finish).await;
367    }
368    let HttpClient { endpoint, http } = client;
369    let Channel {
370        rx: outgoing,
371        tx: incoming,
372    } = channel;
373    let mut outgoing = finishable_outgoing(outgoing, finish);
374    let (sse_event_tx, mut sse_event_rx) = mpsc::unbounded::<SseMessage>();
375    let connection = HttpConnection::new(endpoint, http);
376    let mut state = ClientState {
377        connection: connection.clone(),
378        open_session_streams: HashSet::new(),
379        pending_requests: HashMap::new(),
380        incoming,
381    };
382    let mut lifecycle = HttpTransportLifecycle::new(connection);
383    let mut posts = PostQueues::default();
384    let mut buffered_outgoing = VecDeque::new();
385    let mut outgoing_closed = false;
386
387    let result = 'transport: loop {
388        if outgoing_closed && buffered_outgoing.is_empty() && posts.is_empty() {
389            break Ok(());
390        }
391
392        let event = {
393            let outgoing_next = async {
394                if let Some(frame) = buffered_outgoing.pop_front() {
395                    Some(frame)
396                } else if outgoing_closed {
397                    futures::future::pending().await
398                } else {
399                    outgoing.next().await
400                }
401            }
402            .fuse();
403            let sse_event_next = sse_event_rx.next().fuse();
404            let sse_failure_next = lifecycle.next_sse_failure().fuse();
405            let ordered_post_next = posts.ordered.next_completion().fuse();
406            let response_post_next = posts.responses.next_completion().fuse();
407            pin_mut!(
408                outgoing_next,
409                sse_event_next,
410                sse_failure_next,
411                ordered_post_next,
412                response_post_next
413            );
414
415            futures::select! {
416                msg = outgoing_next => HttpLoopEvent::Outgoing(msg),
417                event = sse_event_next => HttpLoopEvent::SseEvent(event),
418                failure = sse_failure_next => HttpLoopEvent::SseFailure(failure),
419                post = ordered_post_next => HttpLoopEvent::Post(post),
420                post = response_post_next => HttpLoopEvent::Post(post),
421            }
422        };
423
424        let frame = match event {
425            HttpLoopEvent::Outgoing(msg) => {
426                let Some(frame) = msg else {
427                    outgoing_closed = true;
428                    continue;
429                };
430                frame
431            }
432            HttpLoopEvent::SseEvent(event) => {
433                let Some(event) = event else {
434                    continue;
435                };
436                let open_session_ids = state.sessions_to_open_for_responses(&event.frame);
437                state.deliver_frame(event.frame);
438                for session_id in open_session_ids {
439                    match lifecycle
440                        .start_sse(
441                            Some(session_id),
442                            sse_event_tx.clone(),
443                            SseStartContext {
444                                events: &mut sse_event_rx,
445                                outgoing: &mut outgoing,
446                                buffered_outgoing: &mut buffered_outgoing,
447                                posts: &mut posts,
448                                state: &mut state,
449                            },
450                        )
451                        .await
452                    {
453                        Ok(SseStartOutcome::Established) => {}
454                        Ok(SseStartOutcome::OutgoingClosed)
455                            if buffered_outgoing.is_empty() && posts.is_empty() =>
456                        {
457                            break 'transport Ok(());
458                        }
459                        Ok(SseStartOutcome::OutgoingClosed) => {
460                            break 'transport Err(sse_setup_blocked_output_error());
461                        }
462                        Err(error) => break 'transport Err(error),
463                    }
464                }
465                continue;
466            }
467            HttpLoopEvent::SseFailure(failure) => {
468                break Err(sse_failure_error(failure));
469            }
470            HttpLoopEvent::Post(completed) => {
471                if let Err(error) = handle_completed_post(&mut state, completed) {
472                    break Err(error);
473                }
474                continue;
475            }
476        };
477
478        let is_response_only = is_response_only_frame(&frame);
479        let msg = match frame {
480            TransportFrame::Single(message) => message,
481            frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => {
482                if state.connection.connection_id().is_none() {
483                    break Err(AcpError::invalid_request()
484                        .data("ACP HTTP transport: first message must be `initialize`"));
485                }
486                match state.prepare_frame_post(frame) {
487                    // Response-only batches answer SSE-delivered callbacks and
488                    // must not be blocked behind the request they answer.
489                    Ok((post, session_ids)) => {
490                        for session_id in session_ids {
491                            match lifecycle
492                                .start_sse(
493                                    Some(session_id),
494                                    sse_event_tx.clone(),
495                                    SseStartContext {
496                                        events: &mut sse_event_rx,
497                                        outgoing: &mut outgoing,
498                                        buffered_outgoing: &mut buffered_outgoing,
499                                        posts: &mut posts,
500                                        state: &mut state,
501                                    },
502                                )
503                                .await
504                            {
505                                Ok(SseStartOutcome::Established) => {}
506                                Ok(SseStartOutcome::OutgoingClosed) => {
507                                    break 'transport Err(sse_setup_blocked_output_error());
508                                }
509                                Err(error) => break 'transport Err(error),
510                            }
511                        }
512                        if is_response_only {
513                            posts.responses.push(post);
514                        } else {
515                            posts.ordered.push(post);
516                        }
517                    }
518                    Err(error) => {
519                        error!("POST failed");
520                        break Err(AcpError::internal_error().data(format!("POST: {error}")));
521                    }
522                }
523                continue;
524            }
525        };
526
527        if state.connection.connection_id().is_none() {
528            if !is_initialize_request(&msg) {
529                break Err(AcpError::invalid_request()
530                    .data("ACP HTTP transport: first message must be `initialize`"));
531            }
532            match state.initialize(msg).await {
533                Ok(InitializeOutcome::Connected) => {
534                    match lifecycle
535                        .start_sse(
536                            None,
537                            sse_event_tx.clone(),
538                            SseStartContext {
539                                events: &mut sse_event_rx,
540                                outgoing: &mut outgoing,
541                                buffered_outgoing: &mut buffered_outgoing,
542                                posts: &mut posts,
543                                state: &mut state,
544                            },
545                        )
546                        .await
547                    {
548                        Ok(SseStartOutcome::Established) => {}
549                        Ok(SseStartOutcome::OutgoingClosed) if buffered_outgoing.is_empty() => {
550                            break 'transport Ok(());
551                        }
552                        Ok(SseStartOutcome::OutgoingClosed) => {
553                            break 'transport Err(sse_setup_blocked_output_error());
554                        }
555                        Err(error) => break 'transport Err(error),
556                    }
557                }
558                Ok(InitializeOutcome::Rejected) => {}
559                Err(e) => {
560                    error!("initialize failed");
561                    break Err(AcpError::internal_error().data(format!("initialize: {e}")));
562                }
563            }
564            continue;
565        }
566
567        if let Some(session_id) = session_id_from_message(&msg) {
568            for session_id in state.register_session_streams([session_id]) {
569                match lifecycle
570                    .start_sse(
571                        Some(session_id),
572                        sse_event_tx.clone(),
573                        SseStartContext {
574                            events: &mut sse_event_rx,
575                            outgoing: &mut outgoing,
576                            buffered_outgoing: &mut buffered_outgoing,
577                            posts: &mut posts,
578                            state: &mut state,
579                        },
580                    )
581                    .await
582                {
583                    Ok(SseStartOutcome::Established) => {}
584                    Ok(SseStartOutcome::OutgoingClosed) => {
585                        break 'transport Err(sse_setup_blocked_output_error());
586                    }
587                    Err(error) => break 'transport Err(error),
588                }
589            }
590        }
591
592        match state.prepare_post(msg) {
593            // Responses answer SSE-delivered callbacks and must not be blocked
594            // behind a POST that may be waiting for that callback response.
595            Ok(post) if is_response_only => posts.responses.push(post),
596            Ok(post) => posts.ordered.push(post),
597            Err(e) => {
598                error!("POST failed");
599                break Err(AcpError::internal_error().data(format!("POST: {e}")));
600            }
601        }
602    };
603
604    lifecycle.close().await;
605    result
606}
607
608fn sse_failure_error(failure: SseFailure) -> AcpError {
609    let scope = failure.session_id.as_deref().unwrap_or("connection");
610    error!(
611        session_scoped = failure.session_id.is_some(),
612        "SSE stream ended"
613    );
614    AcpError::internal_error().data(format!("{scope} SSE stream ended: {}", failure.error))
615}
616
617fn sse_setup_blocked_output_error() -> AcpError {
618    AcpError::internal_error()
619        .data("outgoing channel closed while accepted messages awaited SSE stream establishment")
620}
621
622fn handle_completed_post(
623    state: &mut ClientState,
624    completed: CompletedPost,
625) -> Result<(), AcpError> {
626    let CompletedPost {
627        pending_requests,
628        result,
629    } = completed;
630    if let Err(error) = result {
631        state.remove_pending_requests(&pending_requests);
632        error!("POST failed");
633        Err(AcpError::internal_error().data(format!("POST: {error}")))
634    } else {
635        Ok(())
636    }
637}
638
639fn queue_response_post(
640    state: &mut ClientState,
641    posts: &mut PostQueues,
642    frame: TransportFrame,
643) -> Result<(), AcpError> {
644    let post = match frame {
645        TransportFrame::Single(message) => state.prepare_post(message),
646        frame @ (TransportFrame::Malformed { .. } | TransportFrame::Batch(_)) => {
647            state.prepare_frame_post(frame).map(|(post, session_ids)| {
648                debug_assert_eq!(session_ids, Vec::<String>::new());
649                post
650            })
651        }
652    }
653    .map_err(|error| {
654        error!("POST failed");
655        AcpError::internal_error().data(format!("POST: {error}"))
656    })?;
657    posts.responses.push(post);
658    Ok(())
659}
660
661fn is_response_only_frame(frame: &TransportFrame) -> bool {
662    match frame {
663        TransportFrame::Single(RawJsonRpcMessage::Response(_)) => true,
664        TransportFrame::Batch(batch) => batch.entries().all(|entry| match entry {
665            TransportBatchEntry::Message(RawJsonRpcMessage::Response(_)) => true,
666            TransportBatchEntry::Malformed { raw, .. } => is_response_only_shape(raw),
667            TransportBatchEntry::Message(
668                RawJsonRpcMessage::Request(_) | RawJsonRpcMessage::Notification(_),
669            ) => false,
670        }),
671        TransportFrame::Malformed { raw, .. } => {
672            serde_json::from_str(raw).is_ok_and(|value| is_response_only_shape(&value))
673        }
674        TransportFrame::Single(
675            RawJsonRpcMessage::Request(_) | RawJsonRpcMessage::Notification(_),
676        ) => false,
677    }
678}
679
680enum HttpLoopEvent {
681    Outgoing(Option<TransportFrame>),
682    SseEvent(Option<SseMessage>),
683    SseFailure(SseFailure),
684    Post(CompletedPost),
685}
686
687#[derive(Debug)]
688struct SseFailure {
689    session_id: Option<String>,
690    error: String,
691}
692
693#[derive(Debug)]
694struct SseMessage {
695    frame: TransportFrame,
696}
697
698#[derive(Clone, Debug)]
699struct HttpConnection {
700    endpoint: url::Url,
701    http: reqwest::Client,
702    connection_id: Arc<StdMutex<Option<String>>>,
703}
704
705impl HttpConnection {
706    fn new(endpoint: url::Url, http: reqwest::Client) -> Self {
707        Self {
708            endpoint,
709            http,
710            connection_id: Arc::new(StdMutex::new(None)),
711        }
712    }
713
714    fn post(&self) -> reqwest::RequestBuilder {
715        self.http.post(self.endpoint.clone())
716    }
717
718    fn get(&self) -> reqwest::RequestBuilder {
719        self.http.get(self.endpoint.clone())
720    }
721
722    fn set_connection_id(&self, connection_id: String) {
723        *self.connection_id.lock().expect("mutex poisoned") = Some(connection_id);
724    }
725
726    fn connection_id(&self) -> Option<String> {
727        self.connection_id.lock().expect("mutex poisoned").clone()
728    }
729
730    fn take_connection_id(&self) -> Option<String> {
731        self.connection_id.lock().expect("mutex poisoned").take()
732    }
733
734    fn clear_connection_id(&self, expected: &str) {
735        let mut connection_id = self.connection_id.lock().expect("mutex poisoned");
736        if connection_id.as_deref() == Some(expected) {
737            *connection_id = None;
738        }
739    }
740
741    async fn close(&self) {
742        let Some(connection_id) = self.connection_id() else {
743            return;
744        };
745        Self::send_close(
746            self.http.clone(),
747            self.endpoint.clone(),
748            connection_id.clone(),
749        )
750        .await;
751        self.clear_connection_id(&connection_id);
752    }
753
754    fn spawn_close(&self) {
755        let Some(connection_id) = self.take_connection_id() else {
756            return;
757        };
758        let http = self.http.clone();
759        let endpoint = self.endpoint.clone();
760        match tokio::runtime::Handle::try_current() {
761            Ok(handle) => {
762                drop(handle.spawn(Self::send_close(http, endpoint, connection_id)));
763            }
764            Err(_) => {
765                debug!("failed to spawn HTTP DELETE");
766            }
767        }
768    }
769
770    async fn send_close(http: reqwest::Client, endpoint: url::Url, connection_id: String) {
771        if http
772            .delete(endpoint)
773            .header(HEADER_CONNECTION_ID, connection_id)
774            .send()
775            .await
776            .is_err()
777        {
778            debug!("DELETE failed (ignored)");
779        }
780    }
781}
782
783#[derive(Debug)]
784struct HttpTransportLifecycle {
785    connection: HttpConnection,
786    sse_tasks: SseTasks,
787}
788
789#[derive(Clone, Copy, Debug, Eq, PartialEq)]
790enum SseStartOutcome {
791    Established,
792    OutgoingClosed,
793}
794
795struct SseStartContext<'a> {
796    events: &'a mut mpsc::UnboundedReceiver<SseMessage>,
797    outgoing: &'a mut (dyn Stream<Item = TransportFrame> + Unpin + Send),
798    buffered_outgoing: &'a mut VecDeque<TransportFrame>,
799    posts: &'a mut PostQueues,
800    state: &'a mut ClientState,
801}
802
803impl HttpTransportLifecycle {
804    fn new(connection: HttpConnection) -> Self {
805        Self {
806            connection,
807            sse_tasks: SseTasks::default(),
808        }
809    }
810
811    async fn start_sse(
812        &mut self,
813        session_id: Option<String>,
814        event_tx: UnboundedSender<SseMessage>,
815        context: SseStartContext<'_>,
816    ) -> Result<SseStartOutcome, AcpError> {
817        let SseStartContext {
818            events,
819            outgoing,
820            buffered_outgoing,
821            posts,
822            state,
823        } = context;
824        let mut establishing = FuturesUnordered::new();
825        establishing.push(self.begin_sse(session_id, event_tx.clone()));
826
827        loop {
828            if establishing.is_empty() {
829                return Ok(SseStartOutcome::Established);
830            }
831            let outcome = {
832                let failure = self.sse_tasks.next_failure().fuse();
833                let established_next = establishing.next().fuse();
834                let sse_event_next = events.next().fuse();
835                let outgoing_next = outgoing.next().fuse();
836                let ordered_post_next = posts.ordered.next_completion().fuse();
837                let response_post_next = posts.responses.next_completion().fuse();
838                pin_mut!(
839                    failure,
840                    established_next,
841                    sse_event_next,
842                    outgoing_next,
843                    ordered_post_next,
844                    response_post_next
845                );
846                futures::select_biased! {
847                    failure = failure => SseStartWait::Failure(failure),
848                    established = established_next => SseStartWait::Established(established),
849                    event = sse_event_next => SseStartWait::SseEvent(event),
850                    post = response_post_next => SseStartWait::Post(post),
851                    post = ordered_post_next => SseStartWait::Post(post),
852                    outgoing = outgoing_next => SseStartWait::Outgoing(outgoing),
853                }
854            };
855            match outcome {
856                SseStartWait::Established(Some(Ok(()))) => {}
857                SseStartWait::Established(Some(Err(_))) => {
858                    return Err(sse_failure_error(self.sse_tasks.next_failure().await));
859                }
860                SseStartWait::Established(None) => {
861                    return Ok(SseStartOutcome::Established);
862                }
863                SseStartWait::Failure(failure) => return Err(sse_failure_error(failure)),
864                SseStartWait::SseEvent(Some(event)) => {
865                    let open_session_ids = state.sessions_to_open_for_responses(&event.frame);
866                    state.deliver_frame(event.frame);
867                    for session_id in open_session_ids {
868                        establishing.push(self.begin_sse(Some(session_id), event_tx.clone()));
869                    }
870                }
871                SseStartWait::SseEvent(None) => {
872                    return Err(AcpError::internal_error().data("SSE event channel closed"));
873                }
874                SseStartWait::Post(completed) => handle_completed_post(state, completed)?,
875                SseStartWait::Outgoing(Some(frame)) if is_response_only_frame(&frame) => {
876                    queue_response_post(state, posts, frame)?;
877                }
878                SseStartWait::Outgoing(Some(frame)) => buffered_outgoing.push_back(frame),
879                SseStartWait::Outgoing(None) => return Ok(SseStartOutcome::OutgoingClosed),
880            }
881        }
882    }
883
884    fn begin_sse(
885        &mut self,
886        session_id: Option<String>,
887        event_tx: UnboundedSender<SseMessage>,
888    ) -> futures::channel::oneshot::Receiver<()> {
889        let (established_tx, established_rx) = futures::channel::oneshot::channel();
890        self.sse_tasks.push(run_sse(
891            self.connection.clone(),
892            session_id,
893            event_tx,
894            established_tx,
895        ));
896        established_rx
897    }
898
899    async fn next_sse_failure(&mut self) -> SseFailure {
900        self.sse_tasks.next_failure().await
901    }
902
903    async fn close(&mut self) {
904        self.connection.close().await;
905        self.sse_tasks.abort_all();
906    }
907}
908
909enum SseStartWait {
910    Established(Option<Result<(), futures::channel::oneshot::Canceled>>),
911    Failure(SseFailure),
912    SseEvent(Option<SseMessage>),
913    Post(CompletedPost),
914    Outgoing(Option<TransportFrame>),
915}
916
917impl Drop for HttpTransportLifecycle {
918    fn drop(&mut self) {
919        self.sse_tasks.abort_all();
920        self.connection.spawn_close();
921    }
922}
923
924fn run_sse(
925    connection: HttpConnection,
926    session_id: Option<String>,
927    event_tx: UnboundedSender<SseMessage>,
928    established_tx: futures::channel::oneshot::Sender<()>,
929) -> BoxFuture<'static, SseFailure> {
930    Box::pin(async move {
931        let label = session_id.clone();
932        let error = match read_sse(connection, session_id, event_tx, established_tx).await {
933            Ok(()) => "SSE stream closed".to_string(),
934            Err(e) => e,
935        };
936        warn!(session_scoped = label.is_some(), "SSE stream ended");
937        SseFailure {
938            session_id: label,
939            error,
940        }
941    })
942}
943
944#[derive(Debug, Default)]
945struct SseTasks {
946    handles: FuturesUnordered<BoxFuture<'static, SseFailure>>,
947}
948
949impl SseTasks {
950    fn push(&mut self, task: BoxFuture<'static, SseFailure>) {
951        self.handles.push(task);
952    }
953
954    async fn next_failure(&mut self) -> SseFailure {
955        loop {
956            if let Some(failure) = self.handles.next().await {
957                return failure;
958            }
959            futures::future::pending::<()>().await;
960        }
961    }
962
963    fn abort_all(&mut self) {
964        self.handles = FuturesUnordered::new();
965    }
966}
967
968struct ClientState {
969    connection: HttpConnection,
970    open_session_streams: HashSet<String>,
971    pending_requests: HashMap<RequestId, VecDeque<String>>,
972    incoming: futures::channel::mpsc::UnboundedSender<TransportFrame>,
973}
974
975struct PendingPost {
976    pending_requests: Vec<(RequestId, String)>,
977    response: BoxFuture<'static, Result<(), String>>,
978}
979
980impl PendingPost {
981    fn into_completion(self) -> BoxFuture<'static, CompletedPost> {
982        let Self {
983            pending_requests,
984            response,
985        } = self;
986        async move {
987            CompletedPost {
988                pending_requests,
989                result: response.await,
990            }
991        }
992        .boxed()
993    }
994}
995
996#[derive(Debug)]
997struct CompletedPost {
998    pending_requests: Vec<(RequestId, String)>,
999    result: Result<(), String>,
1000}
1001
1002#[derive(Default)]
1003struct PostQueue {
1004    queued: VecDeque<PendingPost>,
1005    in_flight: Option<BoxFuture<'static, CompletedPost>>,
1006}
1007
1008#[derive(Default)]
1009struct PostQueues {
1010    ordered: PostQueue,
1011    responses: PostQueue,
1012}
1013
1014impl PostQueues {
1015    fn is_empty(&self) -> bool {
1016        self.ordered.is_empty() && self.responses.is_empty()
1017    }
1018}
1019
1020impl PostQueue {
1021    fn push(&mut self, post: PendingPost) {
1022        self.queued.push_back(post);
1023        self.start_next();
1024    }
1025
1026    async fn next_completion(&mut self) -> CompletedPost {
1027        loop {
1028            self.start_next();
1029            if let Some(in_flight) = self.in_flight.as_mut() {
1030                let completed = in_flight.await;
1031                self.in_flight = None;
1032                return completed;
1033            }
1034            futures::future::pending::<()>().await;
1035        }
1036    }
1037
1038    fn start_next(&mut self) {
1039        if self.in_flight.is_none()
1040            && let Some(post) = self.queued.pop_front()
1041        {
1042            self.in_flight = Some(post.into_completion());
1043        }
1044    }
1045
1046    fn is_empty(&self) -> bool {
1047        self.queued.is_empty() && self.in_flight.is_none()
1048    }
1049}
1050
1051#[derive(Clone, Copy, Debug, Eq, PartialEq)]
1052enum InitializeOutcome {
1053    Connected,
1054    Rejected,
1055}
1056
1057impl ClientState {
1058    async fn initialize(&self, msg: RawJsonRpcMessage) -> Result<InitializeOutcome, String> {
1059        let response = self
1060            .connection
1061            .post()
1062            .header("Content-Type", "application/json")
1063            .header("Accept", "application/json")
1064            .json(&msg)
1065            .send()
1066            .await
1067            .map_err(|e| e.to_string())?;
1068
1069        let connection_id = response
1070            .headers()
1071            .get(HEADER_CONNECTION_ID)
1072            .and_then(|v| v.to_str().ok())
1073            .map(String::from);
1074        if let Some(connection_id) = &connection_id {
1075            self.connection.set_connection_id(connection_id.clone());
1076        }
1077
1078        if !response.status().is_success() {
1079            let status = response.status();
1080            let body = response.text().await.unwrap_or_default();
1081            return Err(format!("HTTP {status}: {body}"));
1082        }
1083
1084        let body = response.text().await.map_err(|error| error.to_string())?;
1085        let message = match TransportFrame::parse_json(&body) {
1086            TransportFrame::Single(message) => message,
1087            TransportFrame::Malformed { error, .. } => {
1088                return Err(format!("invalid initialize response: {error}"));
1089            }
1090            TransportFrame::Batch(_) => {
1091                return Err("initialize response must not be a JSON-RPC batch".to_string());
1092            }
1093        };
1094
1095        if matches!(
1096            message,
1097            RawJsonRpcMessage::Response(RpcResponse::Error { .. })
1098        ) {
1099            self.deliver(message);
1100            self.connection.close().await;
1101            return Ok(InitializeOutcome::Rejected);
1102        }
1103
1104        connection_id
1105            .ok_or_else(|| format!("server did not return {HEADER_CONNECTION_ID} header"))?;
1106        self.deliver(message);
1107        Ok(InitializeOutcome::Connected)
1108    }
1109
1110    fn prepare_post(&mut self, msg: RawJsonRpcMessage) -> Result<PendingPost, String> {
1111        let session_id = validated_session_id(&msg)?;
1112        let connection_id = self
1113            .connection
1114            .connection_id()
1115            .ok_or_else(|| "POST attempted before initialize".to_string())?;
1116        let mut request = self
1117            .connection
1118            .post()
1119            .header("Accept", "application/json")
1120            .header(HEADER_CONNECTION_ID, connection_id)
1121            .json(&msg);
1122        if let Some(session_id) = session_id {
1123            request = request.header(HEADER_SESSION_ID, session_id);
1124        }
1125
1126        let pending_requests = pending_request_for_message(&msg)
1127            .into_iter()
1128            .collect::<Vec<_>>();
1129        self.track_pending_requests(&pending_requests);
1130
1131        let response = async move {
1132            let response = request.send().await.map_err(|e| e.to_string())?;
1133            if response.status().as_u16() != 202 && !response.status().is_success() {
1134                let status = response.status();
1135                let body = response.text().await.unwrap_or_default();
1136                return Err(format!("HTTP {status}: {body}"));
1137            }
1138            Ok(())
1139        };
1140        Ok(PendingPost {
1141            pending_requests,
1142            response: response.boxed(),
1143        })
1144    }
1145
1146    fn prepare_frame_post(
1147        &mut self,
1148        frame: TransportFrame,
1149    ) -> Result<(PendingPost, Vec<String>), String> {
1150        let bookkeeping = FrameBookkeeping::for_frame(&frame)?;
1151        let connection_id = self
1152            .connection
1153            .connection_id()
1154            .ok_or_else(|| "POST attempted before initialize".to_string())?;
1155        let body = frame.to_json().map_err(|error| error.to_string())?;
1156        let request = self
1157            .connection
1158            .post()
1159            .header("Content-Type", "application/json")
1160            .header("Accept", "application/json")
1161            .header(HEADER_CONNECTION_ID, connection_id)
1162            .body(body);
1163        let response = async move {
1164            let response = request.send().await.map_err(|error| error.to_string())?;
1165            if response.status().as_u16() != 202 && !response.status().is_success() {
1166                let status = response.status();
1167                let body = response.text().await.unwrap_or_default();
1168                return Err(format!("HTTP {status}: {body}"));
1169            }
1170            Ok(())
1171        };
1172        self.track_pending_requests(&bookkeeping.pending_requests);
1173        let session_ids = self.register_session_streams(bookkeeping.session_ids);
1174        Ok((
1175            PendingPost {
1176                pending_requests: bookkeeping.pending_requests,
1177                response: response.boxed(),
1178            },
1179            session_ids,
1180        ))
1181    }
1182
1183    fn track_pending_requests(&mut self, pending_requests: &[(RequestId, String)]) {
1184        for (id, method) in pending_requests {
1185            self.pending_requests
1186                .entry(id.clone())
1187                .or_default()
1188                .push_back(method.clone());
1189        }
1190    }
1191
1192    fn remove_pending_requests(&mut self, pending_requests: &[(RequestId, String)]) {
1193        for (id, method) in pending_requests.iter().rev() {
1194            let remove_entry = self.pending_requests.get_mut(id).is_some_and(|methods| {
1195                if let Some(index) = methods.iter().rposition(|candidate| candidate == method) {
1196                    methods.remove(index);
1197                }
1198                methods.is_empty()
1199            });
1200            if remove_entry {
1201                self.pending_requests.remove(id);
1202            }
1203        }
1204    }
1205
1206    fn take_pending_request_method(&mut self, id: &RequestId) -> Option<String> {
1207        let (method, remove_entry) = {
1208            let methods = self.pending_requests.get_mut(id)?;
1209            (methods.pop_front(), methods.is_empty())
1210        };
1211        if remove_entry {
1212            self.pending_requests.remove(id);
1213        }
1214        method
1215    }
1216
1217    fn register_session_streams(
1218        &mut self,
1219        session_ids: impl IntoIterator<Item = String>,
1220    ) -> Vec<String> {
1221        session_ids
1222            .into_iter()
1223            .filter(|session_id| self.open_session_streams.insert(session_id.clone()))
1224            .collect()
1225    }
1226
1227    fn sessions_to_open_for_responses(&mut self, frame: &TransportFrame) -> Vec<String> {
1228        match frame {
1229            TransportFrame::Single(message) => self
1230                .session_to_open_for_response(message)
1231                .into_iter()
1232                .collect(),
1233            TransportFrame::Batch(batch) => batch
1234                .entries()
1235                .filter_map(|entry| match entry {
1236                    TransportBatchEntry::Message(message) => {
1237                        self.session_to_open_for_response(message)
1238                    }
1239                    TransportBatchEntry::Malformed { .. } => None,
1240                })
1241                .collect(),
1242            TransportFrame::Malformed { .. } => Vec::new(),
1243        }
1244    }
1245
1246    fn session_to_open_for_response(&mut self, msg: &RawJsonRpcMessage) -> Option<String> {
1247        let RawJsonRpcMessage::Response(response) = msg else {
1248            return None;
1249        };
1250        let id = msg.response_id().and_then(pending_request_key)?;
1251        let method = self.take_pending_request_method(&id);
1252
1253        if !method.as_deref().is_some_and(is_session_opening_method) {
1254            return None;
1255        }
1256        let RpcResponse::Result { result, .. } = response else {
1257            return None;
1258        };
1259        let session_id = result
1260            .get("sessionId")
1261            .and_then(|v| v.as_str())
1262            .map(String::from)?;
1263
1264        if self.open_session_streams.insert(session_id.clone()) {
1265            Some(session_id)
1266        } else {
1267            None
1268        }
1269    }
1270
1271    fn deliver(&self, msg: RawJsonRpcMessage) {
1272        self.deliver_frame(TransportFrame::Single(msg));
1273    }
1274
1275    fn deliver_frame(&self, frame: TransportFrame) {
1276        if self.incoming.unbounded_send(frame).is_err() {
1277            debug!("upstream channel closed; dropping inbound message");
1278        }
1279    }
1280}
1281
1282#[derive(Default)]
1283struct FrameBookkeeping {
1284    session_ids: Vec<String>,
1285    pending_requests: Vec<(RequestId, String)>,
1286}
1287
1288impl FrameBookkeeping {
1289    fn for_frame(frame: &TransportFrame) -> Result<Self, String> {
1290        let mut bookkeeping = Self::default();
1291        match frame {
1292            TransportFrame::Single(message) => bookkeeping.add_message(message)?,
1293            TransportFrame::Batch(batch) => {
1294                for entry in batch.entries() {
1295                    if let TransportBatchEntry::Message(message) = entry {
1296                        bookkeeping.add_message(message)?;
1297                    }
1298                }
1299            }
1300            TransportFrame::Malformed { .. } => {}
1301        }
1302        Ok(bookkeeping)
1303    }
1304
1305    fn add_message(&mut self, message: &RawJsonRpcMessage) -> Result<(), String> {
1306        if let Some(session_id) = validated_session_id(message)?
1307            && !self.session_ids.contains(&session_id)
1308        {
1309            self.session_ids.push(session_id);
1310        }
1311        if let Some(pending_request) = pending_request_for_message(message) {
1312            self.pending_requests.push(pending_request);
1313        }
1314        Ok(())
1315    }
1316}
1317
1318fn validated_session_id(msg: &RawJsonRpcMessage) -> Result<Option<String>, String> {
1319    let Some(method) = method_for_message(msg) else {
1320        return Ok(None);
1321    };
1322    let session_id = session_id_from_message(msg);
1323    if method_requires_session_header(method) && session_id.is_none() {
1324        return Err(format!("method `{method}` requires sessionId in params"));
1325    }
1326    Ok(session_id)
1327}
1328
1329fn is_session_opening_method(method: &str) -> bool {
1330    matches!(method, "session/new" | "session/fork")
1331}
1332
1333async fn read_sse(
1334    connection: HttpConnection,
1335    session_id: Option<String>,
1336    event_tx: UnboundedSender<SseMessage>,
1337    established_tx: futures::channel::oneshot::Sender<()>,
1338) -> Result<(), String> {
1339    let connection_id = connection
1340        .connection_id()
1341        .ok_or_else(|| "SSE attempted before initialize".to_string())?;
1342    let mut request = connection
1343        .get()
1344        .header("Accept", "text/event-stream")
1345        .header(HEADER_CONNECTION_ID, connection_id);
1346    if let Some(session_id) = &session_id {
1347        request = request.header(HEADER_SESSION_ID, session_id);
1348    }
1349
1350    let response = request.send().await.map_err(|e| e.to_string())?;
1351    if !response.status().is_success() {
1352        return Err(format!("HTTP {}", response.status()));
1353    }
1354    trace!(session_scoped = session_id.is_some(), "SSE stream open");
1355    let _ = established_tx.send(());
1356
1357    let mut events = eventsource_stream::EventStream::new(response.bytes_stream());
1358    while let Some(event) = events.next().await {
1359        let event = event.map_err(|e| e.to_string())?;
1360        let payload = event.data;
1361        if payload.is_empty() {
1362            continue;
1363        }
1364        let frame = TransportFrame::parse_json(&payload);
1365
1366        if event_tx.unbounded_send(SseMessage { frame }).is_err() {
1367            return Err("upstream channel closed".to_string());
1368        }
1369    }
1370    Ok(())
1371}
1372
1373fn pending_request_for_message(msg: &RawJsonRpcMessage) -> Option<(RequestId, String)> {
1374    let RawJsonRpcMessage::Request(request) = msg else {
1375        return None;
1376    };
1377    pending_request_key(&request.id).map(|id| (id, request.method.to_string()))
1378}
1379
1380fn pending_request_key(id: &RequestId) -> Option<RequestId> {
1381    match id {
1382        RequestId::Null => None,
1383        RequestId::Number(_) | RequestId::Str(_) => Some(id.clone()),
1384    }
1385}
1386
1387async fn run_ws(
1388    client: HttpClient,
1389    channel: Channel,
1390    finish: Option<oneshot::Receiver<()>>,
1391) -> Result<(), AcpError> {
1392    let HttpClient { endpoint, http } = client;
1393
1394    let (ws_stream, status) = connect_ws(&http, endpoint).await?;
1395    trace!(status = %status, "WebSocket connection established");
1396    let (ws_tx, ws_rx) = ws_stream.split();
1397
1398    drive_ws_with_finish(ws_tx, ws_rx, channel, finish).await
1399}
1400
1401fn websocket_http_url(mut endpoint: url::Url) -> Result<url::Url, AcpError> {
1402    let scheme = match endpoint.scheme() {
1403        "ws" => "http",
1404        "wss" => "https",
1405        other => {
1406            return Err(
1407                AcpError::internal_error().data(format!("unsupported WebSocket scheme: {other}"))
1408            );
1409        }
1410    };
1411    endpoint
1412        .set_scheme(scheme)
1413        .map_err(|()| AcpError::internal_error().data("failed to convert WebSocket URL"))?;
1414    Ok(endpoint)
1415}
1416
1417async fn connect_ws(
1418    http: &reqwest::Client,
1419    endpoint: url::Url,
1420) -> Result<
1421    (
1422        async_tungstenite::WebSocketStream<
1423            async_tungstenite::tokio::TokioAdapter<reqwest::Upgraded>,
1424        >,
1425        reqwest::StatusCode,
1426    ),
1427    AcpError,
1428> {
1429    let http_url = websocket_http_url(endpoint)?;
1430    let key = async_tungstenite::tungstenite::handshake::client::generate_key();
1431    let expected_accept =
1432        async_tungstenite::tungstenite::handshake::derive_accept_key(key.as_bytes());
1433
1434    // Explicit headers override client defaults, but reqwest can overwrite them
1435    // later with non-tunnel Proxy::headers. See configure_http's caller constraint.
1436    let response = http
1437        .get(http_url)
1438        .version(reqwest::Version::HTTP_11)
1439        .header("Connection", "Upgrade")
1440        .header("Upgrade", "websocket")
1441        .header("Sec-WebSocket-Version", "13")
1442        .header("Sec-WebSocket-Key", &key)
1443        .send()
1444        .await
1445        .map_err(|e| AcpError::internal_error().data(format!("WebSocket connect failed: {e}")))?;
1446    let status = response.status();
1447    validate_ws_response(
1448        response.version(),
1449        status,
1450        response.headers(),
1451        &expected_accept,
1452    )?;
1453
1454    let upgraded = response
1455        .upgrade()
1456        .await
1457        .map_err(|e| AcpError::internal_error().data(format!("WebSocket connect failed: {e}")))?;
1458    let ws_stream = async_tungstenite::WebSocketStream::from_raw_socket(
1459        async_tungstenite::tokio::TokioAdapter::new(upgraded),
1460        async_tungstenite::tungstenite::protocol::Role::Client,
1461        None,
1462    )
1463    .await;
1464    Ok((ws_stream, status))
1465}
1466
1467fn validate_ws_response(
1468    version: reqwest::Version,
1469    status: reqwest::StatusCode,
1470    headers: &reqwest::header::HeaderMap,
1471    expected_accept: &str,
1472) -> Result<(), AcpError> {
1473    let invalid =
1474        |reason| AcpError::internal_error().data(format!("WebSocket connect failed: {reason}"));
1475    if version != reqwest::Version::HTTP_11 {
1476        return Err(invalid(format!(
1477            "expected HTTP/1.1, received {version:?}; preconfigured TLS must use HTTP/1.1 ALPN"
1478        )));
1479    }
1480    if status != reqwest::StatusCode::SWITCHING_PROTOCOLS {
1481        return Err(invalid(format!("unexpected status {status}")));
1482    }
1483    let mut upgrades = headers.get_all("upgrade").iter();
1484    if !upgrades
1485        .next()
1486        .is_some_and(|value| value.as_bytes().eq_ignore_ascii_case(b"websocket"))
1487        || upgrades.next().is_some()
1488    {
1489        return Err(invalid("invalid upgrade header".to_string()));
1490    }
1491    let connection_upgrade = headers.get_all("connection").iter().any(|value| {
1492        value.to_str().is_ok_and(|value| {
1493            value.split(',').any(|part| {
1494                part.trim_matches([' ', '\t'])
1495                    .eq_ignore_ascii_case("upgrade")
1496            })
1497        })
1498    });
1499    if !connection_upgrade {
1500        return Err(invalid("invalid connection header".to_string()));
1501    }
1502    let mut accepts = headers.get_all("sec-websocket-accept").iter();
1503    if accepts.next().map(reqwest::header::HeaderValue::as_bytes)
1504        != Some(expected_accept.as_bytes())
1505        || accepts.next().is_some()
1506    {
1507        return Err(invalid("invalid Sec-WebSocket-Accept".to_string()));
1508    }
1509    // ACP does not negotiate subprotocols or extensions. In particular, passing
1510    // an extension through to from_raw_socket does not enable support for it
1511    // (e.g. compression).
1512    for header in ["sec-websocket-protocol", "sec-websocket-extensions"] {
1513        if headers.contains_key(header) {
1514            return Err(invalid(format!("unsupported {header}")));
1515        }
1516    }
1517    Ok(())
1518}
1519
1520trait WsSink {
1521    fn send(
1522        &mut self,
1523        message: WsMessage,
1524    ) -> impl std::future::Future<Output = Result<(), String>> + Send;
1525}
1526
1527impl<S> WsSink for async_tungstenite::WebSocketSender<S>
1528where
1529    S: futures::AsyncRead + futures::AsyncWrite + Unpin + Send,
1530{
1531    async fn send(&mut self, message: WsMessage) -> Result<(), String> {
1532        async_tungstenite::WebSocketSender::send(self, message)
1533            .await
1534            .map_err(|error| error.to_string())
1535    }
1536}
1537
1538#[cfg(test)]
1539async fn drive_ws<Tx, Rx, RxError>(ws_tx: Tx, ws_rx: Rx, channel: Channel) -> Result<(), AcpError>
1540where
1541    Tx: WsSink,
1542    Rx: Stream<Item = Result<WsMessage, RxError>> + Unpin,
1543    RxError: std::fmt::Display,
1544{
1545    drive_ws_with_finish(ws_tx, ws_rx, channel, None).await
1546}
1547
1548async fn drive_ws_with_finish<Tx, Rx, RxError>(
1549    mut ws_tx: Tx,
1550    mut ws_rx: Rx,
1551    channel: Channel,
1552    finish: Option<oneshot::Receiver<()>>,
1553) -> Result<(), AcpError>
1554where
1555    Tx: WsSink,
1556    Rx: Stream<Item = Result<WsMessage, RxError>> + Unpin,
1557    RxError: std::fmt::Display,
1558{
1559    let Channel {
1560        rx: outgoing,
1561        tx: incoming,
1562    } = channel;
1563    let mut outgoing = finishable_outgoing(outgoing, finish);
1564    let writer = async move {
1565        while let Some(frame) = outgoing.next().await {
1566            let text = match frame.to_json() {
1567                Ok(text) => text,
1568                Err(error) => {
1569                    error!("failed to serialize outbound frame");
1570                    return Err(AcpError::internal_error().data(format!("serialize: {error}")));
1571                }
1572            };
1573            if let Err(error) = ws_tx.send(WsMessage::Text(text.into())).await {
1574                error!("WebSocket send failed");
1575                return Err(AcpError::internal_error().data(format!("ws send: {error}")));
1576            }
1577        }
1578
1579        ws_tx
1580            .send(WsMessage::Close(None))
1581            .await
1582            .map_err(|error| AcpError::internal_error().data(format!("ws close: {error}")))?;
1583        Ok(())
1584    };
1585
1586    let reader = async move {
1587        let mut discard_incoming = false;
1588        loop {
1589            match ws_rx.next().await {
1590                Some(Ok(WsMessage::Text(text))) => {
1591                    if discard_incoming {
1592                        continue;
1593                    }
1594                    let frame = TransportFrame::parse_json(text.as_str());
1595                    if incoming.unbounded_send(frame).is_err() {
1596                        debug!(
1597                            "upstream channel closed; discarding WS input while draining output"
1598                        );
1599                        discard_incoming = true;
1600                    }
1601                }
1602                Some(Ok(WsMessage::Binary(_))) => {
1603                    warn!("ignoring binary WebSocket frame (ACP uses text)");
1604                }
1605                Some(Ok(WsMessage::Ping(_) | WsMessage::Pong(_) | WsMessage::Frame(_))) => {}
1606                Some(Ok(WsMessage::Close(frame))) => {
1607                    debug!("server closed WebSocket");
1608                    return Err(AcpError::internal_error()
1609                        .data(format!("WebSocket closed by peer: {frame:?}")));
1610                }
1611                Some(Err(e)) => {
1612                    error!("WebSocket receive failed");
1613                    return Err(AcpError::internal_error().data(format!("ws recv: {e}")));
1614                }
1615                None => {
1616                    return Err(AcpError::internal_error().data("WebSocket stream ended"));
1617                }
1618            }
1619        }
1620    };
1621
1622    pin_mut!(writer, reader);
1623    match futures::future::select(writer, reader).await {
1624        futures::future::Either::Left((result, _))
1625        | futures::future::Either::Right((result, _)) => result,
1626    }
1627}
1628
1629#[cfg(test)]
1630mod tests {
1631    use std::{
1632        convert::Infallible,
1633        sync::{
1634            Arc,
1635            atomic::{AtomicBool, AtomicUsize, Ordering},
1636        },
1637        time::Duration,
1638    };
1639
1640    use agent_client_protocol::{TransportBatch, UntypedMessage, schema::v1::RequestId};
1641    use axum::{
1642        Json, Router,
1643        extract::{WebSocketUpgrade, ws::Message as AxumWsMessage},
1644        http::{HeaderMap, HeaderValue, StatusCode},
1645        response::{IntoResponse, Sse, sse::Event},
1646        routing::{get, post},
1647    };
1648    use serde_json::json;
1649    use tokio::{
1650        net::TcpListener,
1651        sync::Notify,
1652        time::{sleep, timeout},
1653    };
1654
1655    use super::*;
1656
1657    struct PostsThenExitClient {
1658        finish: Arc<Notify>,
1659        finished: Arc<Notify>,
1660        escaped_tx: futures::channel::oneshot::Sender<
1661            futures::channel::mpsc::UnboundedSender<TransportFrame>,
1662        >,
1663    }
1664
1665    struct InitializeThenExitClient {
1666        sse_started: Arc<Notify>,
1667        finished: Arc<Notify>,
1668    }
1669
1670    struct QueueOutgoingThenText {
1671        text: Option<WsMessage>,
1672        outgoing: Option<mpsc::UnboundedSender<TransportFrame>>,
1673    }
1674
1675    struct RecordingWsSink(mpsc::UnboundedSender<WsMessage>);
1676
1677    struct BackpressuredWsSink {
1678        output: mpsc::UnboundedSender<WsMessage>,
1679        started: mpsc::UnboundedSender<()>,
1680        release: Option<futures::channel::oneshot::Receiver<()>>,
1681    }
1682
1683    struct ReleaseBackpressureOnPoll {
1684        started: mpsc::UnboundedReceiver<()>,
1685        release: Option<futures::channel::oneshot::Sender<()>>,
1686    }
1687
1688    fn single_frame(message: RawJsonRpcMessage) -> TransportFrame {
1689        TransportFrame::Single(message)
1690    }
1691
1692    #[tokio::test]
1693    async fn finish_seals_escaped_senders_and_drains_accepted_frames() {
1694        let (tx, rx) = mpsc::unbounded();
1695        let escaped = tx.clone();
1696        let (finish_tx, finish_rx) = oneshot::channel();
1697        for method in ["custom/first", "custom/second"] {
1698            tx.unbounded_send(single_frame(
1699                RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(),
1700            ))
1701            .unwrap();
1702        }
1703        let mut outgoing = finishable_outgoing(rx, Some(finish_rx));
1704        finish_tx.send(()).unwrap();
1705        for method in ["custom/first", "custom/second"] {
1706            let message = into_single_message(outgoing.next().await.unwrap()).unwrap();
1707            assert_eq!(method_for_message(&message), Some(method));
1708            assert!(escaped.is_closed());
1709        }
1710        assert!(outgoing.next().await.is_none());
1711        assert!(
1712            escaped
1713                .unbounded_send(single_frame(
1714                    RawJsonRpcMessage::notification("custom/too-late".to_string(), json!({}))
1715                        .unwrap(),
1716                ))
1717                .is_err()
1718        );
1719    }
1720
1721    #[tokio::test]
1722    async fn dropping_finish_signal_does_not_close_outgoing() {
1723        let (tx, rx) = mpsc::unbounded();
1724        let (finish_tx, finish_rx) = oneshot::channel();
1725        let mut outgoing = finishable_outgoing(rx, Some(finish_rx));
1726        drop(finish_tx);
1727        assert!(outgoing.next().now_or_never().is_none());
1728        assert!(!tx.is_closed());
1729        tx.unbounded_send(single_frame(
1730            RawJsonRpcMessage::notification("custom/still-open".to_string(), json!({})).unwrap(),
1731        ))
1732        .unwrap();
1733        assert!(outgoing.next().await.is_some());
1734        drop(tx);
1735        assert!(outgoing.next().await.is_none());
1736    }
1737
1738    #[tokio::test]
1739    async fn converted_driver_supports_finish_and_natural_producer_eof() {
1740        for request_finish in [false, true] {
1741            let client = HttpClient::new("http://127.0.0.1:1").unwrap();
1742            let (caller, driver) = ConnectTo::<Client>::into_channel_and_future(client);
1743            let mut driver = driver.unwrap();
1744            if request_finish {
1745                assert!(driver.request_finish());
1746            } else {
1747                drop(caller.tx);
1748            }
1749            timeout(Duration::from_secs(1), driver)
1750                .await
1751                .unwrap()
1752                .unwrap();
1753        }
1754    }
1755
1756    fn into_single_message(frame: TransportFrame) -> Result<RawJsonRpcMessage, AcpError> {
1757        match frame {
1758            TransportFrame::Single(message) => Ok(message),
1759            TransportFrame::Malformed { error, .. } => Err(error),
1760            TransportFrame::Batch(_) => {
1761                Err(AcpError::internal_error().data("expected one JSON-RPC message"))
1762            }
1763        }
1764    }
1765
1766    trait TransportFrameTestExt {
1767        fn unwrap(self) -> RawJsonRpcMessage;
1768    }
1769
1770    impl TransportFrameTestExt for TransportFrame {
1771        fn unwrap(self) -> RawJsonRpcMessage {
1772            into_single_message(self).unwrap()
1773        }
1774    }
1775
1776    #[test]
1777    fn malformed_response_shapes_bypass_only_when_the_whole_frame_is_response_only() {
1778        let standalone_response = TransportFrame::parse_json(
1779            r#"{"jsonrpc":"2.0","id":1,"result":{},"error":{"code":-32603}}"#,
1780        );
1781        assert!(is_response_only_frame(&standalone_response));
1782
1783        let response_batch = TransportFrame::parse_json(
1784            r#"[
1785                {"jsonrpc":"2.0","id":1,"result":{}},
1786                {"jsonrpc":"2.0","id":2,"result":{},"error":{"code":-32603}}
1787            ]"#,
1788        );
1789        assert!(is_response_only_frame(&response_batch));
1790
1791        let scalar_batch = TransportFrame::parse_json(
1792            r#"[
1793                {"jsonrpc":"2.0","id":1,"result":{}},
1794                17
1795            ]"#,
1796        );
1797        assert!(!is_response_only_frame(&scalar_batch));
1798
1799        let call_shaped = TransportFrame::parse_json(
1800            r#"{"jsonrpc":"2.0","id":1,"method":"custom/call","result":{}}"#,
1801        );
1802        assert!(!is_response_only_frame(&call_shaped));
1803    }
1804
1805    fn initialized_client_state() -> ClientState {
1806        let connection = HttpConnection::new(
1807            url::Url::parse("http://127.0.0.1/acp").unwrap(),
1808            reqwest::Client::new(),
1809        );
1810        connection.set_connection_id("connection-1".to_string());
1811        let (incoming, _incoming_rx) = mpsc::unbounded();
1812        ClientState {
1813            connection,
1814            open_session_streams: HashSet::new(),
1815            pending_requests: HashMap::new(),
1816            incoming,
1817        }
1818    }
1819
1820    #[test]
1821    fn batch_post_validation_happens_before_tracking_requests_or_sessions() {
1822        let mut state = initialized_client_state();
1823        let frame = TransportFrame::Batch(
1824            TransportBatch::from_messages([
1825                RawJsonRpcMessage::request(
1826                    "custom/valid".to_string(),
1827                    json!({}),
1828                    RequestId::Number(1),
1829                )
1830                .unwrap(),
1831                RawJsonRpcMessage::request(
1832                    "session/prompt".to_string(),
1833                    json!({ "prompt": [] }),
1834                    RequestId::Number(2),
1835                )
1836                .unwrap(),
1837            ])
1838            .unwrap(),
1839        );
1840
1841        let Err(error) = state.prepare_frame_post(frame) else {
1842            panic!("batch should require sessionId for session/prompt");
1843        };
1844
1845        assert_eq!(
1846            error,
1847            "method `session/prompt` requires sessionId in params"
1848        );
1849        assert!(state.pending_requests.is_empty());
1850        assert!(state.open_session_streams.is_empty());
1851    }
1852
1853    #[test]
1854    fn batch_post_tracks_every_non_null_request_and_rolls_back_from_the_back() {
1855        let mut state = initialized_client_state();
1856        state.track_pending_requests(&[(RequestId::Number(7), "session/fork".to_string())]);
1857        let frame = TransportFrame::Batch(
1858            TransportBatch::from_messages([
1859                RawJsonRpcMessage::request(
1860                    "session/fork".to_string(),
1861                    json!({ "sessionId": "source-a" }),
1862                    RequestId::Number(7),
1863                )
1864                .unwrap(),
1865                RawJsonRpcMessage::request(
1866                    "custom/request".to_string(),
1867                    json!({ "sessionId": "source-b" }),
1868                    RequestId::Number(7),
1869                )
1870                .unwrap(),
1871                RawJsonRpcMessage::request(
1872                    "session/fork".to_string(),
1873                    json!({ "sessionId": "source-a" }),
1874                    RequestId::Null,
1875                )
1876                .unwrap(),
1877            ])
1878            .unwrap(),
1879        );
1880
1881        let (post, session_ids) = state.prepare_frame_post(frame).unwrap();
1882
1883        assert_eq!(session_ids, ["source-a", "source-b"]);
1884        assert_eq!(
1885            state.pending_requests.get(&RequestId::Number(7)).unwrap(),
1886            &VecDeque::from([
1887                "session/fork".to_string(),
1888                "session/fork".to_string(),
1889                "custom/request".to_string(),
1890            ])
1891        );
1892        assert_eq!(
1893            post.pending_requests,
1894            [
1895                (RequestId::Number(7), "session/fork".to_string()),
1896                (RequestId::Number(7), "custom/request".to_string()),
1897            ]
1898        );
1899        assert!(!state.pending_requests.contains_key(&RequestId::Null));
1900
1901        state.remove_pending_requests(&post.pending_requests);
1902
1903        assert_eq!(
1904            state.pending_requests.get(&RequestId::Number(7)).unwrap(),
1905            &VecDeque::from(["session/fork".to_string()])
1906        );
1907    }
1908
1909    impl WsSink for RecordingWsSink {
1910        fn send(
1911            &mut self,
1912            message: WsMessage,
1913        ) -> impl std::future::Future<Output = Result<(), String>> + Send {
1914            std::future::ready(
1915                self.0
1916                    .unbounded_send(message)
1917                    .map_err(|error| error.to_string()),
1918            )
1919        }
1920    }
1921
1922    impl WsSink for BackpressuredWsSink {
1923        async fn send(&mut self, message: WsMessage) -> Result<(), String> {
1924            self.output
1925                .unbounded_send(message)
1926                .map_err(|error| error.to_string())?;
1927            if let Some(release) = self.release.take() {
1928                self.started
1929                    .unbounded_send(())
1930                    .map_err(|error| error.to_string())?;
1931                release
1932                    .await
1933                    .map_err(|_| "mock WebSocket reader did not release send".to_string())?;
1934            }
1935            Ok(())
1936        }
1937    }
1938
1939    impl Stream for QueueOutgoingThenText {
1940        type Item = Result<WsMessage, std::io::Error>;
1941
1942        fn poll_next(
1943            mut self: std::pin::Pin<&mut Self>,
1944            _cx: &mut std::task::Context<'_>,
1945        ) -> std::task::Poll<Option<Self::Item>> {
1946            // Make input ready immediately after queueing output. If the output
1947            // branch was polled first it was still empty, so either poll order
1948            // deterministically selects this input frame first.
1949            if let Some(outgoing) = self.outgoing.take() {
1950                for method in ["custom/first", "custom/second"] {
1951                    outgoing
1952                        .unbounded_send(single_frame(
1953                            RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(),
1954                        ))
1955                        .unwrap();
1956                }
1957            }
1958            if let Some(text) = self.text.take() {
1959                return std::task::Poll::Ready(Some(Ok(text)));
1960            }
1961            std::task::Poll::Pending
1962        }
1963    }
1964
1965    impl Stream for ReleaseBackpressureOnPoll {
1966        type Item = Result<WsMessage, std::io::Error>;
1967
1968        fn poll_next(
1969            mut self: std::pin::Pin<&mut Self>,
1970            cx: &mut std::task::Context<'_>,
1971        ) -> std::task::Poll<Option<Self::Item>> {
1972            if let std::task::Poll::Ready(Some(())) =
1973                std::pin::Pin::new(&mut self.started).poll_next(cx)
1974                && let Some(release) = self.release.take()
1975            {
1976                let _result = release.send(());
1977            }
1978            std::task::Poll::Pending
1979        }
1980    }
1981
1982    impl ConnectTo<Agent> for PostsThenExitClient {
1983        async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), AcpError> {
1984            let Self {
1985                finish,
1986                finished,
1987                escaped_tx,
1988            } = self;
1989            let (mut channel, transport) = agent.into_channel_and_future();
1990            let client = async move {
1991                escaped_tx.send(channel.tx.clone()).map_err(|_| {
1992                    AcpError::internal_error().data("escaped sender observer dropped")
1993                })?;
1994                channel
1995                    .tx
1996                    .unbounded_send(single_frame(
1997                        RawJsonRpcMessage::request(
1998                            "initialize".to_string(),
1999                            json!({}),
2000                            RequestId::Number(1),
2001                        )
2002                        .unwrap(),
2003                    ))
2004                    .map_err(|e| {
2005                        AcpError::internal_error().data(format!("send initialize: {e}"))
2006                    })?;
2007                into_single_message(channel.rx.next().await.ok_or_else(|| {
2008                    AcpError::internal_error().data("initialize response channel closed")
2009                })?)?;
2010
2011                for method in ["custom/first", "custom/second"] {
2012                    channel
2013                        .tx
2014                        .unbounded_send(single_frame(
2015                            RawJsonRpcMessage::notification(method.to_string(), json!({})).unwrap(),
2016                        ))
2017                        .map_err(|e| {
2018                            AcpError::internal_error().data(format!("send {method}: {e}"))
2019                        })?;
2020                }
2021
2022                finish.notified().await;
2023                finished.notify_one();
2024                Ok(())
2025            };
2026
2027            let transport = async move {
2028                if let Some(transport) = transport {
2029                    transport.await?;
2030                }
2031                Ok::<(), AcpError>(())
2032            };
2033            let ((), ()) = futures::try_join!(transport, client)?;
2034            Ok(())
2035        }
2036    }
2037
2038    impl ConnectTo<Agent> for InitializeThenExitClient {
2039        async fn connect_to(self, agent: impl ConnectTo<Client>) -> Result<(), AcpError> {
2040            let Self {
2041                sse_started,
2042                finished,
2043            } = self;
2044            let (mut channel, transport) = agent.into_channel_and_future();
2045            let client = async move {
2046                channel
2047                    .tx
2048                    .unbounded_send(single_frame(
2049                        RawJsonRpcMessage::request(
2050                            "initialize".to_string(),
2051                            json!({}),
2052                            RequestId::Number(1),
2053                        )
2054                        .unwrap(),
2055                    ))
2056                    .map_err(|error| {
2057                        AcpError::internal_error().data(format!("send initialize: {error}"))
2058                    })?;
2059                into_single_message(channel.rx.next().await.ok_or_else(|| {
2060                    AcpError::internal_error().data("initialize response channel closed")
2061                })?)?;
2062
2063                sse_started.notified().await;
2064                finished.notify_one();
2065                Ok(())
2066            };
2067
2068            let transport = async move {
2069                if let Some(transport) = transport {
2070                    transport.await?;
2071                }
2072                Ok::<(), AcpError>(())
2073            };
2074            let ((), ()) = futures::try_join!(transport, client)?;
2075            Ok(())
2076        }
2077    }
2078
2079    #[test]
2080    fn new_targets_standard_acp_endpoint() {
2081        assert_eq!(
2082            HttpClient::new("http://example.com")
2083                .unwrap()
2084                .endpoint
2085                .as_str(),
2086            "http://example.com/acp"
2087        );
2088        assert_eq!(
2089            HttpClient::new("http://example.com/proxy")
2090                .unwrap()
2091                .endpoint
2092                .as_str(),
2093            "http://example.com/proxy/acp"
2094        );
2095        assert_eq!(
2096            HttpClient::new("http://example.com/proxy/acp")
2097                .unwrap()
2098                .endpoint
2099                .as_str(),
2100            "http://example.com/proxy/acp"
2101        );
2102    }
2103
2104    #[test]
2105    fn with_endpoint_preserves_explicit_endpoint_path() {
2106        assert_eq!(
2107            HttpClient::with_endpoint("http://example.com/agent")
2108                .unwrap()
2109                .endpoint
2110                .as_str(),
2111            "http://example.com/agent"
2112        );
2113        assert_eq!(
2114            HttpClient::builder_with_endpoint("ws://example.com/custom/acp?token=abc")
2115                .build()
2116                .unwrap()
2117                .endpoint
2118                .as_str(),
2119            "ws://example.com/custom/acp?token=abc"
2120        );
2121    }
2122
2123    #[test]
2124    fn builder_uses_the_same_base_url_rule_for_all_transports() {
2125        for scheme in ["http", "https", "ws", "wss"] {
2126            for (path, expected) in [
2127                ("", "/acp"),
2128                ("/", "/acp"),
2129                ("/proxy/", "/proxy/acp"),
2130                ("/proxy/acp/", "/proxy/acp"),
2131                ("/proxy/acp/nested", "/proxy/acp/nested/acp"),
2132            ] {
2133                let url = format!("{scheme}://example.com{path}?key=value");
2134                let client = HttpClient::builder(&url).build().unwrap();
2135                assert_eq!(
2136                    client.endpoint.as_str(),
2137                    format!("{scheme}://example.com{expected}?key=value")
2138                );
2139                let exact = HttpClient::builder_with_endpoint(&url).build().unwrap();
2140                assert_eq!(exact.endpoint, url::Url::parse(&url).unwrap());
2141            }
2142        }
2143    }
2144
2145    #[test]
2146    fn constructors_reject_invalid_urls_and_unsupported_schemes() {
2147        for build in [HttpClient::builder, HttpClient::builder_with_endpoint] {
2148            assert!(matches!(
2149                build("not a URL".to_string()).build(),
2150                Err(HttpClientError::InvalidUrl(_))
2151            ));
2152            for scheme in ["file", "ftp", "custom"] {
2153                assert!(matches!(
2154                    build(format!("{scheme}://example.com/acp")).build(),
2155                    Err(HttpClientError::UnsupportedScheme(actual)) if actual == scheme
2156                ));
2157            }
2158        }
2159    }
2160
2161    #[test]
2162    fn prebuilt_http_client_requires_an_http_endpoint() {
2163        let http = reqwest::Client::new();
2164        for scheme in ["http", "https"] {
2165            let endpoint = format!("{scheme}://example.com/custom?key=value");
2166            let client = HttpClient::from_http_client(&endpoint, http.clone()).unwrap();
2167            assert_eq!(client.endpoint.as_str(), endpoint);
2168        }
2169        for scheme in ["ws", "wss"] {
2170            assert!(matches!(
2171                HttpClient::from_http_client(format!("{scheme}://example.com/acp"), http.clone()),
2172                Err(HttpClientError::WebSocketRequiresBuilder)
2173            ));
2174        }
2175        assert!(matches!(
2176            HttpClient::from_http_client("ftp://example.com/acp", http),
2177            Err(HttpClientError::UnsupportedScheme(_))
2178        ));
2179    }
2180
2181    #[test]
2182    #[allow(deprecated)]
2183    fn deprecated_constructors_preserve_http_paths_and_reject_websockets() {
2184        let http = reqwest::Client::new();
2185        for scheme in ["http", "https"] {
2186            for path in ["", "/proxy", "/proxy/acp/"] {
2187                let url = format!("{scheme}://example.com{path}?key=value");
2188                let legacy = HttpClient::with_client(&url, http.clone()).unwrap();
2189                assert_eq!(legacy.endpoint, HttpClient::new(&url).unwrap().endpoint);
2190                let exact = HttpClient::with_endpoint_and_client(&url, http.clone()).unwrap();
2191                assert_eq!(
2192                    exact.endpoint,
2193                    HttpClient::with_endpoint(&url).unwrap().endpoint
2194                );
2195            }
2196        }
2197        for scheme in ["ws", "wss"] {
2198            let url = format!("{scheme}://example.com/custom");
2199            assert!(matches!(
2200                HttpClient::with_client(&url, http.clone()),
2201                Err(HttpClientError::WebSocketRequiresBuilder)
2202            ));
2203            assert!(matches!(
2204                HttpClient::with_endpoint_and_client(&url, http.clone()),
2205                Err(HttpClientError::WebSocketRequiresBuilder)
2206            ));
2207        }
2208    }
2209
2210    #[test]
2211    fn builder_propagates_http_configuration_errors_without_panicking() {
2212        let error = HttpClient::builder("ws://example.com")
2213            .configure_http(|http| http.user_agent("\n"))
2214            .build()
2215            .unwrap_err();
2216        assert!(matches!(error, HttpClientError::Reqwest(_)));
2217    }
2218
2219    #[test]
2220    fn client_and_builder_debug_do_not_expose_default_headers() {
2221        let mut headers = HeaderMap::new();
2222        headers.insert(
2223            "x-api-key",
2224            HeaderValue::from_static("private-header-value"),
2225        );
2226        let builder = HttpClient::builder("ws://example.com")
2227            .configure_http(|http| http.default_headers(headers));
2228        assert!(!format!("{builder:?}").contains("private-header-value"));
2229        let client = builder.build().unwrap();
2230        assert!(!format!("{client:?}").contains("private-header-value"));
2231        assert_eq!(client.clone().endpoint, client.endpoint);
2232    }
2233
2234    #[tokio::test]
2235    async fn post_sends_cancel_request_without_session_header() {
2236        let (capture_tx, mut capture_rx) = tokio::sync::mpsc::unbounded_channel();
2237        let post_count = Arc::new(AtomicUsize::new(0));
2238        let app = Router::new().route(
2239            "/acp",
2240            post({
2241                let capture_tx = capture_tx.clone();
2242                let post_count = post_count.clone();
2243                move |headers: HeaderMap, Json(message): Json<RawJsonRpcMessage>| {
2244                    let capture_tx = capture_tx.clone();
2245                    let post_count = post_count.clone();
2246                    async move {
2247                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2248                            return initialize_response().await.into_response();
2249                        }
2250
2251                        capture_tx
2252                            .send((headers.get(HEADER_SESSION_ID).cloned(), message))
2253                            .unwrap();
2254                        StatusCode::ACCEPTED.into_response()
2255                    }
2256                }
2257            })
2258            .get(pending_sse)
2259            .delete(|| async { StatusCode::ACCEPTED }),
2260        );
2261        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2262        let addr = listener.local_addr().unwrap();
2263        let server = tokio::spawn(async move {
2264            axum::serve(listener, app).await.unwrap();
2265        });
2266        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2267        let (mut caller, transport) = Channel::duplex();
2268        let transport = tokio::spawn(run(client, transport));
2269
2270        caller
2271            .tx
2272            .unbounded_send(single_frame(
2273                RawJsonRpcMessage::request(
2274                    "initialize".to_string(),
2275                    json!({}),
2276                    RequestId::Number(1),
2277                )
2278                .unwrap(),
2279            ))
2280            .unwrap();
2281        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2282            .await
2283            .unwrap()
2284            .unwrap()
2285            .unwrap();
2286        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2287
2288        caller
2289            .tx
2290            .unbounded_send(single_frame(
2291                RawJsonRpcMessage::notification(
2292                    "$/cancel_request".to_string(),
2293                    json!({
2294                        "requestId": 2,
2295                        "sessionId": "session-1"
2296                    }),
2297                )
2298                .unwrap(),
2299            ))
2300            .unwrap();
2301
2302        let (session_header, message) = timeout(Duration::from_secs(1), capture_rx.recv())
2303            .await
2304            .unwrap()
2305            .unwrap();
2306        assert!(session_header.is_none());
2307        assert!(matches!(
2308            message,
2309            RawJsonRpcMessage::Notification(notification)
2310                if notification.method.as_ref() == "$/cancel_request"
2311        ));
2312
2313        drop(caller);
2314        timeout(Duration::from_secs(1), transport)
2315            .await
2316            .unwrap()
2317            .unwrap()
2318            .unwrap();
2319
2320        server.abort();
2321    }
2322
2323    #[tokio::test]
2324    async fn http_preserves_batch_frames_across_post_and_sse() {
2325        let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel();
2326        let post_count = Arc::new(AtomicUsize::new(0));
2327        let emit_sse = Arc::new(Notify::new());
2328        let inbound_batch = json!([
2329            {
2330                "jsonrpc": "2.0",
2331                "method": "custom/inbound-one",
2332                "params": {}
2333            },
2334            {
2335                "jsonrpc": "2.0",
2336                "method": "custom/inbound-two",
2337                "params": {}
2338            }
2339        ]);
2340        let app = Router::new().route(
2341            "/acp",
2342            post({
2343                let post_count = post_count.clone();
2344                move |body: String| {
2345                    let post_count = post_count.clone();
2346                    let post_tx = post_tx.clone();
2347                    async move {
2348                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2349                            return initialize_response().await.into_response();
2350                        }
2351
2352                        post_tx
2353                            .send(serde_json::from_str::<serde_json::Value>(&body).unwrap())
2354                            .unwrap();
2355                        StatusCode::ACCEPTED.into_response()
2356                    }
2357                }
2358            })
2359            .get({
2360                let emit_sse = emit_sse.clone();
2361                let inbound_batch = inbound_batch.clone();
2362                move || {
2363                    let emit_sse = emit_sse.clone();
2364                    let inbound_batch = inbound_batch.clone();
2365                    async move {
2366                        let stream = async_stream::stream! {
2367                            emit_sse.notified().await;
2368                            yield Ok::<_, Infallible>(
2369                                Event::default().data(inbound_batch.to_string()),
2370                            );
2371                            futures::future::pending::<()>().await;
2372                        };
2373                        Sse::new(stream)
2374                    }
2375                }
2376            })
2377            .delete(|| async { StatusCode::ACCEPTED }),
2378        );
2379        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2380        let addr = listener.local_addr().unwrap();
2381        let server = tokio::spawn(async move {
2382            axum::serve(listener, app).await.unwrap();
2383        });
2384        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2385        let (mut caller, transport) = Channel::duplex();
2386        let transport = tokio::spawn(run(client, transport));
2387
2388        caller
2389            .tx
2390            .unbounded_send(single_frame(
2391                RawJsonRpcMessage::request(
2392                    "initialize".to_string(),
2393                    json!({}),
2394                    RequestId::Number(1),
2395                )
2396                .unwrap(),
2397            ))
2398            .unwrap();
2399        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2400            .await
2401            .unwrap()
2402            .unwrap()
2403            .unwrap();
2404        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2405
2406        let outbound_batch = json!([
2407            {
2408                "jsonrpc": "2.0",
2409                "method": "custom/outbound-one",
2410                "params": {}
2411            },
2412            {
2413                "jsonrpc": "2.0",
2414                "method": "custom/outbound-two",
2415                "params": {}
2416            }
2417        ]);
2418        caller
2419            .tx
2420            .unbounded_send(TransportFrame::Batch(
2421                TransportBatch::from_messages([
2422                    RawJsonRpcMessage::notification("custom/outbound-one".to_string(), json!({}))
2423                        .unwrap(),
2424                    RawJsonRpcMessage::notification("custom/outbound-two".to_string(), json!({}))
2425                        .unwrap(),
2426                ])
2427                .unwrap(),
2428            ))
2429            .unwrap();
2430
2431        let posted = timeout(Duration::from_secs(1), post_rx.recv())
2432            .await
2433            .unwrap()
2434            .unwrap();
2435        assert_eq!(posted, outbound_batch);
2436
2437        emit_sse.notify_one();
2438        let inbound = timeout(Duration::from_secs(1), caller.rx.next())
2439            .await
2440            .unwrap()
2441            .unwrap();
2442        assert!(matches!(&inbound, TransportFrame::Batch(_)));
2443        assert_eq!(
2444            serde_json::from_str::<serde_json::Value>(&inbound.to_json().unwrap()).unwrap(),
2445            inbound_batch
2446        );
2447
2448        drop(caller);
2449        timeout(Duration::from_secs(1), transport)
2450            .await
2451            .unwrap()
2452            .unwrap()
2453            .unwrap();
2454
2455        server.abort();
2456    }
2457
2458    #[tokio::test]
2459    async fn batch_fork_opens_source_and_result_session_streams() {
2460        let (post_tx, mut post_rx) = tokio::sync::mpsc::unbounded_channel();
2461        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
2462        let post_count = Arc::new(AtomicUsize::new(0));
2463        let emit_response = Arc::new(Notify::new());
2464        let connection_stream_established = Arc::new(AtomicBool::new(false));
2465        let source_stream_established = Arc::new(AtomicBool::new(false));
2466        let response_batch = json!([
2467            {
2468                "jsonrpc": "2.0",
2469                "id": 2,
2470                "result": { "sessionId": "forked-session" }
2471            }
2472        ]);
2473        let app = Router::new().route(
2474            "/acp",
2475            post({
2476                let post_count = post_count.clone();
2477                let connection_stream_established = connection_stream_established.clone();
2478                let source_stream_established = source_stream_established.clone();
2479                move |body: String| {
2480                    let post_count = post_count.clone();
2481                    let post_tx = post_tx.clone();
2482                    let connection_stream_established = connection_stream_established.clone();
2483                    let source_stream_established = source_stream_established.clone();
2484                    async move {
2485                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2486                            return initialize_response().await.into_response();
2487                        }
2488
2489                        if !connection_stream_established.load(Ordering::SeqCst)
2490                            || !source_stream_established.load(Ordering::SeqCst)
2491                        {
2492                            return StatusCode::CONFLICT.into_response();
2493                        }
2494                        post_tx
2495                            .send(serde_json::from_str::<serde_json::Value>(&body).unwrap())
2496                            .unwrap();
2497                        StatusCode::ACCEPTED.into_response()
2498                    }
2499                }
2500            })
2501            .get({
2502                let emit_response = emit_response.clone();
2503                let response_batch = response_batch.clone();
2504                let connection_stream_established = connection_stream_established.clone();
2505                let source_stream_established = source_stream_established.clone();
2506                move |headers: HeaderMap| {
2507                    let emit_response = emit_response.clone();
2508                    let response_batch = response_batch.clone();
2509                    let get_tx = get_tx.clone();
2510                    let connection_stream_established = connection_stream_established.clone();
2511                    let source_stream_established = source_stream_established.clone();
2512                    async move {
2513                        let session_id = headers
2514                            .get(HEADER_SESSION_ID)
2515                            .and_then(|value| value.to_str().ok())
2516                            .map(String::from);
2517                        let is_connection_stream = session_id.is_none();
2518                        let is_source_stream = session_id.as_deref() == Some("source-session");
2519                        if is_connection_stream {
2520                            sleep(Duration::from_millis(50)).await;
2521                            connection_stream_established.store(true, Ordering::SeqCst);
2522                        }
2523                        if is_source_stream {
2524                            sleep(Duration::from_millis(50)).await;
2525                            source_stream_established.store(true, Ordering::SeqCst);
2526                        }
2527                        get_tx.send(session_id).unwrap();
2528
2529                        let stream = async_stream::stream! {
2530                            if is_source_stream {
2531                                emit_response.notified().await;
2532                                yield Ok::<_, Infallible>(
2533                                    Event::default().data(response_batch.to_string()),
2534                                );
2535                            }
2536                            futures::future::pending::<()>().await;
2537                        };
2538                        Sse::new(stream)
2539                    }
2540                }
2541            })
2542            .delete(|| async { StatusCode::ACCEPTED }),
2543        );
2544        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2545        let addr = listener.local_addr().unwrap();
2546        let server = tokio::spawn(async move {
2547            axum::serve(listener, app).await.unwrap();
2548        });
2549        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2550        let (mut caller, transport) = Channel::duplex();
2551        let transport = tokio::spawn(run(client, transport));
2552
2553        caller
2554            .tx
2555            .unbounded_send(single_frame(
2556                RawJsonRpcMessage::request(
2557                    "initialize".to_string(),
2558                    json!({}),
2559                    RequestId::Number(1),
2560                )
2561                .unwrap(),
2562            ))
2563            .unwrap();
2564        timeout(Duration::from_secs(1), caller.rx.next())
2565            .await
2566            .unwrap()
2567            .unwrap();
2568
2569        caller
2570            .tx
2571            .unbounded_send(TransportFrame::Batch(
2572                TransportBatch::from_messages([RawJsonRpcMessage::request(
2573                    "session/fork".to_string(),
2574                    json!({ "sessionId": "source-session" }),
2575                    RequestId::Number(2),
2576                )
2577                .unwrap()])
2578                .unwrap(),
2579            ))
2580            .unwrap();
2581
2582        let connection_stream = timeout(Duration::from_secs(1), get_rx.recv())
2583            .await
2584            .unwrap()
2585            .unwrap();
2586        assert!(connection_stream.is_none());
2587        let source_stream = timeout(Duration::from_secs(1), get_rx.recv())
2588            .await
2589            .unwrap()
2590            .unwrap();
2591        assert_eq!(source_stream.as_deref(), Some("source-session"));
2592        let posted = timeout(Duration::from_secs(1), post_rx.recv())
2593            .await
2594            .unwrap()
2595            .unwrap();
2596        assert!(posted.is_array(), "outgoing batch must remain an array");
2597
2598        emit_response.notify_one();
2599        let response = timeout(Duration::from_secs(1), caller.rx.next())
2600            .await
2601            .unwrap()
2602            .unwrap();
2603        assert!(matches!(&response, TransportFrame::Batch(_)));
2604        assert_eq!(
2605            serde_json::from_str::<serde_json::Value>(&response.to_json().unwrap()).unwrap(),
2606            response_batch
2607        );
2608        let forked_stream = timeout(Duration::from_secs(1), get_rx.recv())
2609            .await
2610            .unwrap()
2611            .unwrap();
2612        assert_eq!(forked_stream.as_deref(), Some("forked-session"));
2613
2614        drop(caller);
2615        timeout(Duration::from_secs(1), transport)
2616            .await
2617            .unwrap()
2618            .unwrap()
2619            .unwrap();
2620
2621        server.abort();
2622    }
2623
2624    #[tokio::test]
2625    async fn custom_response_with_session_id_does_not_open_session_sse() {
2626        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
2627        let response_ready = Arc::new(tokio::sync::Notify::new());
2628        let post_count = Arc::new(AtomicUsize::new(0));
2629        let app = Router::new().route(
2630            "/acp",
2631            post({
2632                let post_count = post_count.clone();
2633                let response_ready = response_ready.clone();
2634                move |Json(_message): Json<RawJsonRpcMessage>| {
2635                    let post_count = post_count.clone();
2636                    let response_ready = response_ready.clone();
2637                    async move {
2638                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2639                            return initialize_response().await.into_response();
2640                        }
2641
2642                        response_ready.notify_waiters();
2643                        StatusCode::ACCEPTED.into_response()
2644                    }
2645                }
2646            })
2647            .get({
2648                let get_tx = get_tx.clone();
2649                let response_ready = response_ready.clone();
2650                move |headers: HeaderMap| {
2651                    let get_tx = get_tx.clone();
2652                    let response_ready = response_ready.clone();
2653                    async move {
2654                        let session_header = headers
2655                            .get(HEADER_SESSION_ID)
2656                            .and_then(|value| value.to_str().ok())
2657                            .map(String::from);
2658                        get_tx.send(session_header).unwrap();
2659
2660                        let stream = async_stream::stream! {
2661                            response_ready.notified().await;
2662                            yield Ok::<_, Infallible>(sse_event(
2663                                RawJsonRpcMessage::response(
2664                                    RequestId::Number(2),
2665                                    Ok(json!({ "sessionId": "session-1" })),
2666                                ),
2667                            ));
2668                            futures::future::pending::<()>().await;
2669                        };
2670                        Sse::new(stream)
2671                    }
2672                }
2673            })
2674            .delete(|| async { StatusCode::ACCEPTED }),
2675        );
2676        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2677        let addr = listener.local_addr().unwrap();
2678        let server = tokio::spawn(async move {
2679            axum::serve(listener, app).await.unwrap();
2680        });
2681        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2682        let (mut caller, transport) = Channel::duplex();
2683        let transport = tokio::spawn(run(client, transport));
2684
2685        caller
2686            .tx
2687            .unbounded_send(single_frame(
2688                RawJsonRpcMessage::request(
2689                    "initialize".to_string(),
2690                    json!({}),
2691                    RequestId::Number(1),
2692                )
2693                .unwrap(),
2694            ))
2695            .unwrap();
2696        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2697            .await
2698            .unwrap()
2699            .unwrap()
2700            .unwrap();
2701        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2702
2703        let connection_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2704            .await
2705            .unwrap()
2706            .unwrap();
2707        assert!(connection_sse_header.is_none());
2708
2709        caller
2710            .tx
2711            .unbounded_send(single_frame(
2712                RawJsonRpcMessage::request(
2713                    "custom/sessionish".to_string(),
2714                    json!({}),
2715                    RequestId::Number(2),
2716                )
2717                .unwrap(),
2718            ))
2719            .unwrap();
2720        let response = timeout(Duration::from_secs(1), caller.rx.next())
2721            .await
2722            .unwrap()
2723            .unwrap()
2724            .unwrap();
2725        assert!(matches!(
2726            response,
2727            RawJsonRpcMessage::Response(RpcResponse::Result {
2728                id: RequestId::Number(2),
2729                ..
2730            })
2731        ));
2732
2733        assert!(
2734            timeout(Duration::from_millis(100), get_rx.recv())
2735                .await
2736                .is_err(),
2737            "custom response must not open a session SSE stream"
2738        );
2739
2740        drop(caller);
2741        timeout(Duration::from_secs(1), transport)
2742            .await
2743            .unwrap()
2744            .unwrap()
2745            .unwrap();
2746
2747        server.abort();
2748    }
2749
2750    #[tokio::test]
2751    async fn fork_response_with_session_id_opens_session_sse() {
2752        let (get_tx, mut get_rx) = tokio::sync::mpsc::unbounded_channel();
2753        let response_ready = Arc::new(tokio::sync::Notify::new());
2754        let post_count = Arc::new(AtomicUsize::new(0));
2755        let app = Router::new().route(
2756            "/acp",
2757            post({
2758                let post_count = post_count.clone();
2759                let response_ready = response_ready.clone();
2760                move |Json(_message): Json<RawJsonRpcMessage>| {
2761                    let post_count = post_count.clone();
2762                    let response_ready = response_ready.clone();
2763                    async move {
2764                        if post_count.fetch_add(1, Ordering::SeqCst) == 0 {
2765                            return initialize_response().await.into_response();
2766                        }
2767
2768                        response_ready.notify_waiters();
2769                        StatusCode::ACCEPTED.into_response()
2770                    }
2771                }
2772            })
2773            .get({
2774                let get_tx = get_tx.clone();
2775                let response_ready = response_ready.clone();
2776                move |headers: HeaderMap| {
2777                    let get_tx = get_tx.clone();
2778                    let response_ready = response_ready.clone();
2779                    async move {
2780                        let session_header = headers
2781                            .get(HEADER_SESSION_ID)
2782                            .and_then(|value| value.to_str().ok())
2783                            .map(String::from);
2784                        let is_connection_stream = session_header.is_none();
2785                        get_tx.send(session_header).unwrap();
2786
2787                        let stream = async_stream::stream! {
2788                            if is_connection_stream {
2789                                response_ready.notified().await;
2790                                yield Ok::<_, Infallible>(sse_event(
2791                                    RawJsonRpcMessage::response(
2792                                        RequestId::Number(2),
2793                                        Ok(json!({ "sessionId": "forked-session" })),
2794                                    ),
2795                                ));
2796                            }
2797                            futures::future::pending::<()>().await;
2798                        };
2799                        Sse::new(stream)
2800                    }
2801                }
2802            })
2803            .delete(|| async { StatusCode::ACCEPTED }),
2804        );
2805        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2806        let addr = listener.local_addr().unwrap();
2807        let server = tokio::spawn(async move {
2808            axum::serve(listener, app).await.unwrap();
2809        });
2810        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2811        let (mut caller, transport) = Channel::duplex();
2812        let transport = tokio::spawn(run(client, transport));
2813
2814        caller
2815            .tx
2816            .unbounded_send(single_frame(
2817                RawJsonRpcMessage::request(
2818                    "initialize".to_string(),
2819                    json!({}),
2820                    RequestId::Number(1),
2821                )
2822                .unwrap(),
2823            ))
2824            .unwrap();
2825        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
2826            .await
2827            .unwrap()
2828            .unwrap()
2829            .unwrap();
2830        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
2831
2832        let connection_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2833            .await
2834            .unwrap()
2835            .unwrap();
2836        assert!(connection_sse_header.is_none());
2837
2838        caller
2839            .tx
2840            .unbounded_send(single_frame(
2841                RawJsonRpcMessage::request(
2842                    "session/fork".to_string(),
2843                    json!({ "sessionId": "source-session" }),
2844                    RequestId::Number(2),
2845                )
2846                .unwrap(),
2847            ))
2848            .unwrap();
2849        let source_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2850            .await
2851            .unwrap()
2852            .unwrap();
2853        assert_eq!(source_sse_header.as_deref(), Some("source-session"));
2854
2855        let response = timeout(Duration::from_secs(1), caller.rx.next())
2856            .await
2857            .unwrap()
2858            .unwrap()
2859            .unwrap();
2860        assert!(matches!(
2861            response,
2862            RawJsonRpcMessage::Response(RpcResponse::Result {
2863                id: RequestId::Number(2),
2864                ..
2865            })
2866        ));
2867
2868        let fork_sse_header = timeout(Duration::from_secs(1), get_rx.recv())
2869            .await
2870            .unwrap()
2871            .unwrap();
2872        assert_eq!(fork_sse_header.as_deref(), Some("forked-session"));
2873
2874        drop(caller);
2875        timeout(Duration::from_secs(1), transport)
2876            .await
2877            .unwrap()
2878            .unwrap()
2879            .unwrap();
2880
2881        server.abort();
2882    }
2883
2884    #[tokio::test]
2885    async fn only_response_batches_bypass_ordered_posts() {
2886        let slow_started = Arc::new(Notify::new());
2887        let release_slow = Arc::new(Notify::new());
2888        let call_batch_seen = Arc::new(Notify::new());
2889        let response_batch_seen = Arc::new(Notify::new());
2890        let app = Router::new().route(
2891            "/acp",
2892            post({
2893                let slow_started = slow_started.clone();
2894                let release_slow = release_slow.clone();
2895                let call_batch_seen = call_batch_seen.clone();
2896                let response_batch_seen = response_batch_seen.clone();
2897                move |body: String| {
2898                    let slow_started = slow_started.clone();
2899                    let release_slow = release_slow.clone();
2900                    let call_batch_seen = call_batch_seen.clone();
2901                    let response_batch_seen = response_batch_seen.clone();
2902                    async move {
2903                        let value = serde_json::from_str::<serde_json::Value>(&body).unwrap();
2904                        if value.get("method").and_then(serde_json::Value::as_str)
2905                            == Some("initialize")
2906                        {
2907                            return initialize_response().await.into_response();
2908                        }
2909                        if value.get("method").and_then(serde_json::Value::as_str)
2910                            == Some("custom/slow")
2911                        {
2912                            slow_started.notify_one();
2913                            release_slow.notified().await;
2914                        } else if let Some(entries) = value.as_array() {
2915                            if entries.iter().all(|entry| entry.get("method").is_none()) {
2916                                response_batch_seen.notify_one();
2917                            } else {
2918                                call_batch_seen.notify_one();
2919                            }
2920                        }
2921                        StatusCode::ACCEPTED.into_response()
2922                    }
2923                }
2924            })
2925            .get(pending_sse)
2926            .delete(|| async { StatusCode::ACCEPTED }),
2927        );
2928        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
2929        let addr = listener.local_addr().unwrap();
2930        let server = tokio::spawn(async move {
2931            axum::serve(listener, app).await.unwrap();
2932        });
2933        let client = HttpClient::new(format!("http://{addr}")).unwrap();
2934        let (mut caller, transport) = Channel::duplex();
2935        let transport = tokio::spawn(run(client, transport));
2936
2937        caller
2938            .tx
2939            .unbounded_send(single_frame(
2940                RawJsonRpcMessage::request(
2941                    "initialize".to_string(),
2942                    json!({}),
2943                    RequestId::Number(1),
2944                )
2945                .unwrap(),
2946            ))
2947            .unwrap();
2948        timeout(Duration::from_secs(1), caller.rx.next())
2949            .await
2950            .unwrap()
2951            .unwrap();
2952
2953        caller
2954            .tx
2955            .unbounded_send(single_frame(
2956                RawJsonRpcMessage::notification("custom/slow".to_string(), json!({})).unwrap(),
2957            ))
2958            .unwrap();
2959        timeout(Duration::from_secs(1), slow_started.notified())
2960            .await
2961            .unwrap();
2962
2963        caller
2964            .tx
2965            .unbounded_send(TransportFrame::Batch(
2966                TransportBatch::from_messages([
2967                    RawJsonRpcMessage::notification("custom/one".to_string(), json!({})).unwrap(),
2968                    RawJsonRpcMessage::notification("custom/two".to_string(), json!({})).unwrap(),
2969                ])
2970                .unwrap(),
2971            ))
2972            .unwrap();
2973        assert!(
2974            timeout(Duration::from_millis(100), call_batch_seen.notified())
2975                .await
2976                .is_err(),
2977            "call-bearing batches must remain behind an earlier ordered POST"
2978        );
2979
2980        caller
2981            .tx
2982            .unbounded_send(TransportFrame::Batch(
2983                TransportBatch::from_messages([
2984                    RawJsonRpcMessage::response(RequestId::Number(10), Ok(json!({}))),
2985                    RawJsonRpcMessage::response(RequestId::Number(11), Ok(json!({}))),
2986                ])
2987                .unwrap(),
2988            ))
2989            .unwrap();
2990        timeout(Duration::from_secs(1), response_batch_seen.notified())
2991            .await
2992            .expect("response-only batch should bypass the ordered POST queue");
2993
2994        release_slow.notify_one();
2995        timeout(Duration::from_secs(1), call_batch_seen.notified())
2996            .await
2997            .expect("call-bearing batch should be sent after the earlier POST completes");
2998
2999        drop(caller);
3000        timeout(Duration::from_secs(1), transport)
3001            .await
3002            .unwrap()
3003            .unwrap()
3004            .unwrap();
3005
3006        server.abort();
3007    }
3008
3009    #[tokio::test]
3010    async fn client_completion_drains_ordered_posts_in_order() {
3011        let first_started = Arc::new(Notify::new());
3012        let release_first = Arc::new(Notify::new());
3013        let second_seen = Arc::new(Notify::new());
3014        let finish_client = Arc::new(Notify::new());
3015        let client_finished = Arc::new(Notify::new());
3016        let (escaped_tx, escaped_rx) = futures::channel::oneshot::channel();
3017        let app = Router::new().route(
3018            "/acp",
3019            post({
3020                let first_started = first_started.clone();
3021                let release_first = release_first.clone();
3022                let second_seen = second_seen.clone();
3023                move |Json(message): Json<RawJsonRpcMessage>| {
3024                    let first_started = first_started.clone();
3025                    let release_first = release_first.clone();
3026                    let second_seen = second_seen.clone();
3027                    async move {
3028                        if is_initialize_request(&message) {
3029                            return initialize_response().await.into_response();
3030                        }
3031
3032                        match method_for_message(&message) {
3033                            Some("custom/first") => {
3034                                first_started.notify_one();
3035                                release_first.notified().await;
3036                            }
3037                            Some("custom/second") => {
3038                                second_seen.notify_one();
3039                            }
3040                            _ => {}
3041                        }
3042                        StatusCode::ACCEPTED.into_response()
3043                    }
3044                }
3045            })
3046            .get(pending_sse)
3047            .delete(|| async { StatusCode::ACCEPTED }),
3048        );
3049        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3050        let addr = listener.local_addr().unwrap();
3051        let server = tokio::spawn(async move {
3052            axum::serve(listener, app).await.unwrap();
3053        });
3054        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3055        let mut connection = tokio::spawn(client.connect_to(PostsThenExitClient {
3056            finish: finish_client.clone(),
3057            finished: client_finished.clone(),
3058            escaped_tx,
3059        }));
3060        let escaped = timeout(Duration::from_secs(1), escaped_rx)
3061            .await
3062            .unwrap()
3063            .unwrap();
3064
3065        timeout(Duration::from_secs(1), first_started.notified())
3066            .await
3067            .unwrap();
3068        assert!(
3069            timeout(Duration::from_millis(100), second_seen.notified())
3070                .await
3071                .is_err(),
3072            "second POST must not be sent while the first POST is pending"
3073        );
3074
3075        finish_client.notify_one();
3076        timeout(Duration::from_secs(1), client_finished.notified())
3077            .await
3078            .unwrap();
3079        assert!(
3080            timeout(Duration::from_millis(100), &mut connection)
3081                .await
3082                .is_err(),
3083            "HTTP transport returned before its accepted POSTs completed"
3084        );
3085        assert!(
3086            escaped
3087                .unbounded_send(single_frame(
3088                    RawJsonRpcMessage::notification("custom/too-late".to_string(), json!({}),)
3089                        .unwrap()
3090                ))
3091                .is_err(),
3092            "escaped client sender remained open after client completion"
3093        );
3094
3095        release_first.notify_one();
3096        timeout(Duration::from_secs(1), second_seen.notified())
3097            .await
3098            .unwrap();
3099
3100        timeout(Duration::from_secs(1), connection)
3101            .await
3102            .unwrap()
3103            .unwrap()
3104            .unwrap();
3105
3106        server.abort();
3107    }
3108
3109    #[tokio::test]
3110    async fn builder_completion_drains_ordered_posts_and_delete() {
3111        timeout(Duration::from_secs(3), builder_http_finish(false))
3112            .await
3113            .unwrap();
3114    }
3115
3116    #[tokio::test]
3117    async fn builder_completion_preserves_post_failure() {
3118        timeout(Duration::from_secs(3), builder_http_finish(true))
3119            .await
3120            .unwrap();
3121    }
3122
3123    async fn builder_http_finish(fail_first: bool) {
3124        let first_started = Arc::new(Notify::new());
3125        let release_first = Arc::new(Notify::new());
3126        let release_delete = Arc::new(Notify::new());
3127        let (seen_tx, mut seen) = mpsc::unbounded();
3128        let app = Router::new().route(
3129            "/acp",
3130            post({
3131                let first_started = first_started.clone();
3132                let release_first = release_first.clone();
3133                let seen_tx = seen_tx.clone();
3134                move |Json(message): Json<serde_json::Value>| {
3135                    let first_started = first_started.clone();
3136                    let release_first = release_first.clone();
3137                    let seen_tx = seen_tx.clone();
3138                    async move {
3139                        if message["method"] == "initialize" {
3140                            return (
3141                                [(HEADER_CONNECTION_ID, "conn-1")],
3142                                Json(json!({
3143                                    "jsonrpc": "2.0",
3144                                    "id": message["id"],
3145                                    "result": {"protocolVersion": 1, "agentCapabilities": {}}
3146                                })),
3147                            )
3148                                .into_response();
3149                        }
3150                        let method = message["method"].as_str().unwrap().to_string();
3151                        seen_tx.unbounded_send(method.clone()).unwrap();
3152                        if method == "custom/first" {
3153                            first_started.notify_one();
3154                            release_first.notified().await;
3155                            if fail_first {
3156                                return StatusCode::INTERNAL_SERVER_ERROR.into_response();
3157                            }
3158                        }
3159                        StatusCode::ACCEPTED.into_response()
3160                    }
3161                }
3162            })
3163            .get(pending_sse)
3164            .delete({
3165                let release_delete = release_delete.clone();
3166                move || {
3167                    let release_delete = release_delete.clone();
3168                    let seen_tx = seen_tx.clone();
3169                    async move {
3170                        seen_tx.unbounded_send("delete".to_string()).unwrap();
3171                        release_delete.notified().await;
3172                        StatusCode::ACCEPTED
3173                    }
3174                }
3175            }),
3176        );
3177        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3178        let addr = listener.local_addr().unwrap();
3179        let server = tokio::spawn(async move {
3180            axum::serve(listener, app).await.unwrap();
3181        });
3182        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3183        let mut connection = Box::pin(Client.builder().connect_with(client, async move |cx| {
3184            cx.send_request(UntypedMessage::new(
3185                "initialize",
3186                json!({"protocolVersion": 1, "clientCapabilities": {}}),
3187            )?)
3188            .block_task()
3189            .await?;
3190            for method in ["custom/first", "custom/second"] {
3191                cx.send_notification(UntypedMessage::new(method, json!({}))?)?;
3192            }
3193            // Initialization is delivered before connection SSE establishment.
3194            // Finish only once dependent output is flowing on that connection.
3195            first_started.notified().await;
3196            Ok(())
3197        }));
3198
3199        tokio::select! {
3200            result = &mut connection => panic!("returned before first POST drain: {result:?}"),
3201            event = seen.next() => assert_eq!(event.as_deref(), Some("custom/first")),
3202        }
3203        assert!(
3204            connection.as_mut().now_or_never().is_none(),
3205            "a gated POST must keep graceful shutdown pending"
3206        );
3207        assert!(seen.try_recv().is_err(), "ordered POST bypassed its gate");
3208
3209        release_first.notify_one();
3210        if !fail_first {
3211            tokio::select! {
3212                result = &mut connection => panic!("returned before second POST: {result:?}"),
3213                event = seen.next() => assert_eq!(event.as_deref(), Some("custom/second")),
3214            }
3215        }
3216        tokio::select! {
3217            result = &mut connection => panic!("returned before DELETE: {result:?}"),
3218            event = seen.next() => assert_eq!(event.as_deref(), Some("delete")),
3219        }
3220        assert!(
3221            connection.as_mut().now_or_never().is_none(),
3222            "transport must await physical cleanup, not just spawn DELETE"
3223        );
3224        release_delete.notify_one();
3225        let result = connection.await;
3226        if fail_first {
3227            assert!(result.unwrap_err().to_string().contains("500"));
3228        } else {
3229            result.unwrap();
3230        }
3231        assert!(seen.try_recv().is_err());
3232        server.abort();
3233    }
3234
3235    #[tokio::test]
3236    async fn builder_completion_drains_websocket_and_closes_without_peer_eof() {
3237        let release_upgrade = Arc::new(Notify::new());
3238        let (frames_tx, mut frames) = mpsc::unbounded();
3239        let app = Router::new().route(
3240            "/acp",
3241            get({
3242                let release_upgrade = release_upgrade.clone();
3243                move |ws: WebSocketUpgrade| {
3244                    let release_upgrade = release_upgrade.clone();
3245                    let frames_tx = frames_tx.clone();
3246                    async move {
3247                        release_upgrade.notified().await;
3248                        ws.on_upgrade(async move |mut socket| {
3249                            while let Some(Ok(message)) = socket.recv().await {
3250                                let closed = matches!(message, AxumWsMessage::Close(_));
3251                                frames_tx.unbounded_send(message).unwrap();
3252                                if closed {
3253                                    // Do not require the peer to finish its read half.
3254                                    futures::future::pending::<()>().await;
3255                                }
3256                            }
3257                        })
3258                    }
3259                }
3260            }),
3261        );
3262        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3263        let addr = listener.local_addr().unwrap();
3264        let server = tokio::spawn(async move {
3265            axum::serve(listener, app).await.unwrap();
3266        });
3267        let (callback_done_tx, callback_done_rx) = futures::channel::oneshot::channel();
3268        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
3269        let mut connection = Box::pin(Client.builder().connect_with(client, async move |cx| {
3270            for method in ["custom/first", "custom/second"] {
3271                cx.send_notification(UntypedMessage::new(method, json!({}))?)?;
3272            }
3273            callback_done_tx.send(()).unwrap();
3274            Ok(())
3275        }));
3276        assert!(
3277            connection.as_mut().now_or_never().is_none(),
3278            "queued sends must await the gated WebSocket handshake"
3279        );
3280        callback_done_rx.now_or_never().unwrap().unwrap();
3281        release_upgrade.notify_one();
3282        timeout(Duration::from_secs(3), connection)
3283            .await
3284            .unwrap()
3285            .unwrap();
3286        for method in ["custom/first", "custom/second"] {
3287            let frame = timeout(Duration::from_secs(1), frames.next())
3288                .await
3289                .unwrap()
3290                .unwrap();
3291            let AxumWsMessage::Text(text) = frame else {
3292                panic!("expected outbound text frame, got {frame:?}");
3293            };
3294            let message = serde_json::from_str::<RawJsonRpcMessage>(&text).unwrap();
3295            assert_eq!(method_for_message(&message), Some(method));
3296        }
3297        assert!(matches!(
3298            timeout(Duration::from_secs(1), frames.next())
3299                .await
3300                .unwrap(),
3301            Some(AxumWsMessage::Close(None))
3302        ));
3303        server.abort();
3304    }
3305
3306    #[tokio::test]
3307    async fn client_completion_cancels_pending_sse_establishment() {
3308        let sse_started = Arc::new(Notify::new());
3309        let delete_count = Arc::new(AtomicUsize::new(0));
3310        let client_finished = Arc::new(Notify::new());
3311        let app = Router::new().route(
3312            "/acp",
3313            post(initialize_response)
3314                .get({
3315                    let sse_started = sse_started.clone();
3316                    move || {
3317                        let sse_started = sse_started.clone();
3318                        async move {
3319                            sse_started.notify_one();
3320                            futures::future::pending::<StatusCode>().await
3321                        }
3322                    }
3323                })
3324                .delete({
3325                    let delete_count = delete_count.clone();
3326                    move || {
3327                        let delete_count = delete_count.clone();
3328                        async move {
3329                            delete_count.fetch_add(1, Ordering::SeqCst);
3330                            StatusCode::ACCEPTED
3331                        }
3332                    }
3333                }),
3334        );
3335        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3336        let addr = listener.local_addr().unwrap();
3337        let server = tokio::spawn(async move {
3338            axum::serve(listener, app).await.unwrap();
3339        });
3340        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3341        let connection = tokio::spawn(client.connect_to(InitializeThenExitClient {
3342            sse_started,
3343            finished: client_finished.clone(),
3344        }));
3345
3346        timeout(Duration::from_secs(1), client_finished.notified())
3347            .await
3348            .expect("client foreground did not finish after the SSE request started");
3349
3350        timeout(Duration::from_secs(1), connection)
3351            .await
3352            .expect("transport remained blocked on SSE response headers")
3353            .unwrap()
3354            .unwrap();
3355        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3356
3357        server.abort();
3358    }
3359
3360    #[tokio::test]
3361    async fn stalled_sse_establishment_observes_earlier_post_failure() {
3362        let app = Router::new().route(
3363            "/acp",
3364            get(|| async { futures::future::pending::<StatusCode>().await }),
3365        );
3366        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3367        let addr = listener.local_addr().unwrap();
3368        let server = tokio::spawn(async move {
3369            axum::serve(listener, app).await.unwrap();
3370        });
3371
3372        let connection = HttpConnection::new(
3373            url::Url::parse(&format!("http://{addr}/acp")).unwrap(),
3374            reqwest::Client::new(),
3375        );
3376        connection.set_connection_id("connection-1".to_string());
3377        let (incoming, _incoming_rx) = mpsc::unbounded();
3378        let mut state = ClientState {
3379            connection: connection.clone(),
3380            open_session_streams: HashSet::new(),
3381            pending_requests: HashMap::new(),
3382            incoming,
3383        };
3384        let pending_request = (RequestId::Number(7), "custom/earlier".to_string());
3385        state.track_pending_requests(std::slice::from_ref(&pending_request));
3386        let mut posts = PostQueues::default();
3387        posts.ordered.push(PendingPost {
3388            pending_requests: vec![pending_request],
3389            response: async { Err("earlier post failed".to_string()) }.boxed(),
3390        });
3391
3392        let (_outgoing_tx, mut outgoing) = mpsc::unbounded();
3393        let mut buffered_outgoing = VecDeque::new();
3394        let (event_tx, mut event_rx) = mpsc::unbounded();
3395        let mut lifecycle = HttpTransportLifecycle::new(connection);
3396        let error = timeout(
3397            Duration::from_secs(1),
3398            lifecycle.start_sse(
3399                Some("later-session".to_string()),
3400                event_tx,
3401                SseStartContext {
3402                    events: &mut event_rx,
3403                    outgoing: &mut outgoing,
3404                    buffered_outgoing: &mut buffered_outgoing,
3405                    posts: &mut posts,
3406                    state: &mut state,
3407                },
3408            ),
3409        )
3410        .await
3411        .expect("stalled SSE setup hid an earlier POST failure")
3412        .unwrap_err();
3413
3414        assert!(error.to_string().contains("earlier post failed"));
3415        assert!(state.pending_requests.is_empty());
3416
3417        lifecycle.close().await;
3418        server.abort();
3419    }
3420
3421    #[tokio::test]
3422    async fn stalled_sse_establishment_keeps_callback_responses_moving() {
3423        let release_get = Arc::new(Notify::new());
3424        let complete_earlier_post = Arc::new(Notify::new());
3425        let app = Router::new().route(
3426            "/acp",
3427            post({
3428                let release_get = release_get.clone();
3429                let complete_earlier_post = complete_earlier_post.clone();
3430                move || {
3431                    release_get.notify_one();
3432                    complete_earlier_post.notify_one();
3433                    async { StatusCode::ACCEPTED }
3434                }
3435            })
3436            .get({
3437                let release_get = release_get.clone();
3438                move || {
3439                    let release_get = release_get.clone();
3440                    async move {
3441                        release_get.notified().await;
3442                        Sse::new(futures::stream::pending::<Result<Event, Infallible>>())
3443                    }
3444                }
3445            }),
3446        );
3447        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3448        let addr = listener.local_addr().unwrap();
3449        let server = tokio::spawn(async move {
3450            axum::serve(listener, app).await.unwrap();
3451        });
3452
3453        let connection = HttpConnection::new(
3454            url::Url::parse(&format!("http://{addr}/acp")).unwrap(),
3455            reqwest::Client::new(),
3456        );
3457        connection.set_connection_id("connection-1".to_string());
3458        let (incoming, mut incoming_rx) = mpsc::unbounded();
3459        let mut state = ClientState {
3460            connection: connection.clone(),
3461            open_session_streams: HashSet::new(),
3462            pending_requests: HashMap::new(),
3463            incoming,
3464        };
3465        let mut posts = PostQueues::default();
3466        posts.ordered.push(PendingPost {
3467            pending_requests: Vec::new(),
3468            response: async move {
3469                complete_earlier_post.notified().await;
3470                Ok(())
3471            }
3472            .boxed(),
3473        });
3474
3475        let (outgoing_tx, mut outgoing) = mpsc::unbounded();
3476        let outgoing_guard = outgoing_tx.clone();
3477        let mut buffered_outgoing = VecDeque::new();
3478        let (event_tx, mut event_rx) = mpsc::unbounded();
3479        event_tx
3480            .unbounded_send(SseMessage {
3481                frame: single_frame(
3482                    RawJsonRpcMessage::request(
3483                        "test/callback".to_string(),
3484                        json!({}),
3485                        RequestId::Number(99),
3486                    )
3487                    .unwrap(),
3488                ),
3489            })
3490            .unwrap();
3491
3492        let responder = async move {
3493            let callback = incoming_rx
3494                .next()
3495                .await
3496                .expect("callback was not delivered");
3497            assert!(matches!(
3498                into_single_message(callback).unwrap(),
3499                RawJsonRpcMessage::Request(request)
3500                    if request.method.as_ref() == "test/callback"
3501            ));
3502            outgoing_tx
3503                .unbounded_send(single_frame(RawJsonRpcMessage::response(
3504                    RequestId::Number(99),
3505                    Ok(json!({})),
3506                )))
3507                .unwrap();
3508        };
3509        let mut lifecycle = HttpTransportLifecycle::new(connection);
3510        let (outcome, ()) = timeout(Duration::from_secs(1), async {
3511            futures::join!(
3512                lifecycle.start_sse(
3513                    Some("later-session".to_string()),
3514                    event_tx,
3515                    SseStartContext {
3516                        events: &mut event_rx,
3517                        outgoing: &mut outgoing,
3518                        buffered_outgoing: &mut buffered_outgoing,
3519                        posts: &mut posts,
3520                        state: &mut state,
3521                    },
3522                ),
3523                responder,
3524            )
3525        })
3526        .await
3527        .expect("callback response deadlocked behind stalled SSE establishment");
3528
3529        assert_eq!(outcome.unwrap(), SseStartOutcome::Established);
3530        assert!(buffered_outgoing.is_empty());
3531
3532        drop(outgoing_guard);
3533        lifecycle.close().await;
3534        server.abort();
3535    }
3536
3537    #[tokio::test]
3538    async fn pending_sse_establishment_reports_buffered_output_on_shutdown() {
3539        let sse_started = Arc::new(Notify::new());
3540        let delete_count = Arc::new(AtomicUsize::new(0));
3541        let app = Router::new().route(
3542            "/acp",
3543            post(initialize_response)
3544                .get({
3545                    let sse_started = sse_started.clone();
3546                    move || {
3547                        let sse_started = sse_started.clone();
3548                        async move {
3549                            sse_started.notify_one();
3550                            futures::future::pending::<StatusCode>().await
3551                        }
3552                    }
3553                })
3554                .delete({
3555                    let delete_count = delete_count.clone();
3556                    move || {
3557                        let delete_count = delete_count.clone();
3558                        async move {
3559                            delete_count.fetch_add(1, Ordering::SeqCst);
3560                            StatusCode::ACCEPTED
3561                        }
3562                    }
3563                }),
3564        );
3565        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3566        let addr = listener.local_addr().unwrap();
3567        let server = tokio::spawn(async move {
3568            axum::serve(listener, app).await.unwrap();
3569        });
3570        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3571        let (mut caller, transport) = Channel::duplex();
3572        let transport = tokio::spawn(run(client, transport));
3573
3574        caller
3575            .tx
3576            .unbounded_send(single_frame(
3577                RawJsonRpcMessage::request(
3578                    "initialize".to_string(),
3579                    json!({}),
3580                    RequestId::Number(1),
3581                )
3582                .unwrap(),
3583            ))
3584            .unwrap();
3585        timeout(Duration::from_secs(1), caller.rx.next())
3586            .await
3587            .unwrap()
3588            .unwrap();
3589        timeout(Duration::from_secs(1), sse_started.notified())
3590            .await
3591            .expect("connection SSE request did not reach the server");
3592
3593        caller
3594            .tx
3595            .unbounded_send(single_frame(
3596                RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(),
3597            ))
3598            .unwrap();
3599        drop(caller);
3600
3601        let error = timeout(Duration::from_secs(1), transport)
3602            .await
3603            .expect("transport remained blocked on SSE response headers")
3604            .unwrap()
3605            .unwrap_err();
3606        assert!(error.to_string().contains("accepted messages"));
3607        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3608
3609        server.abort();
3610    }
3611
3612    #[tokio::test]
3613    async fn sse_continues_while_post_is_pending() {
3614        let post_started = Arc::new(Notify::new());
3615        let callback_response_seen = Arc::new(Notify::new());
3616        let sse_started = Arc::new(Notify::new());
3617        let (callback_tx, mut callback_rx) = tokio::sync::mpsc::unbounded_channel();
3618        let app = Router::new().route(
3619            "/acp",
3620            post({
3621                let post_started = post_started.clone();
3622                let callback_response_seen = callback_response_seen.clone();
3623                let callback_tx = callback_tx.clone();
3624                move |Json(message): Json<RawJsonRpcMessage>| {
3625                    let post_started = post_started.clone();
3626                    let callback_response_seen = callback_response_seen.clone();
3627                    let callback_tx = callback_tx.clone();
3628                    async move {
3629                        if is_initialize_request(&message) {
3630                            return initialize_response().await.into_response();
3631                        }
3632
3633                        match &message {
3634                            RawJsonRpcMessage::Request(request)
3635                                if request.method.as_ref() == "custom/slow" =>
3636                            {
3637                                post_started.notify_waiters();
3638                                callback_response_seen.notified().await;
3639                                StatusCode::ACCEPTED.into_response()
3640                            }
3641                            RawJsonRpcMessage::Response(
3642                                RpcResponse::Result {
3643                                    id: RequestId::Number(99),
3644                                    ..
3645                                }
3646                                | RpcResponse::Error {
3647                                    id: RequestId::Number(99),
3648                                    ..
3649                                },
3650                            ) => {
3651                                callback_tx.send(message).unwrap();
3652                                callback_response_seen.notify_waiters();
3653                                StatusCode::ACCEPTED.into_response()
3654                            }
3655                            _ => StatusCode::ACCEPTED.into_response(),
3656                        }
3657                    }
3658                }
3659            })
3660            .get({
3661                let post_started = post_started.clone();
3662                let sse_started = sse_started.clone();
3663                move || {
3664                    let post_started = post_started.clone();
3665                    let sse_started = sse_started.clone();
3666                    async move {
3667                        let stream = async_stream::stream! {
3668                            sse_started.notify_waiters();
3669                            post_started.notified().await;
3670                            yield Ok::<_, Infallible>(sse_event(
3671                                RawJsonRpcMessage::request(
3672                                    "client/callback".to_string(),
3673                                    json!({}),
3674                                    RequestId::Number(99),
3675                                )
3676                                .unwrap(),
3677                            ));
3678                            futures::future::pending::<()>().await;
3679                        };
3680                        Sse::new(stream)
3681                    }
3682                }
3683            })
3684            .delete(|| async { StatusCode::ACCEPTED }),
3685        );
3686        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3687        let addr = listener.local_addr().unwrap();
3688        let server = tokio::spawn(async move {
3689            axum::serve(listener, app).await.unwrap();
3690        });
3691        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3692        let (mut caller, transport) = Channel::duplex();
3693        let transport = tokio::spawn(run(client, transport));
3694
3695        caller
3696            .tx
3697            .unbounded_send(single_frame(
3698                RawJsonRpcMessage::request(
3699                    "initialize".to_string(),
3700                    json!({}),
3701                    RequestId::Number(1),
3702                )
3703                .unwrap(),
3704            ))
3705            .unwrap();
3706        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3707            .await
3708            .unwrap()
3709            .unwrap()
3710            .unwrap();
3711        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3712        timeout(Duration::from_secs(1), sse_started.notified())
3713            .await
3714            .unwrap();
3715
3716        caller
3717            .tx
3718            .unbounded_send(single_frame(
3719                RawJsonRpcMessage::request(
3720                    "custom/slow".to_string(),
3721                    json!({}),
3722                    RequestId::Number(2),
3723                )
3724                .unwrap(),
3725            ))
3726            .unwrap();
3727
3728        let callback = timeout(Duration::from_secs(1), caller.rx.next())
3729            .await
3730            .unwrap()
3731            .unwrap()
3732            .unwrap();
3733        assert!(matches!(
3734            callback,
3735            RawJsonRpcMessage::Request(request)
3736                if request.method.as_ref() == "client/callback"
3737                    && request.id == RequestId::Number(99)
3738        ));
3739
3740        caller
3741            .tx
3742            .unbounded_send(single_frame(RawJsonRpcMessage::response(
3743                RequestId::Number(99),
3744                Ok(json!({})),
3745            )))
3746            .unwrap();
3747        let callback_response = timeout(Duration::from_secs(1), callback_rx.recv())
3748            .await
3749            .unwrap()
3750            .unwrap();
3751        assert!(matches!(
3752            callback_response,
3753            RawJsonRpcMessage::Response(RpcResponse::Result {
3754                id: RequestId::Number(99),
3755                ..
3756            })
3757        ));
3758
3759        drop(caller);
3760        timeout(Duration::from_secs(1), transport)
3761            .await
3762            .unwrap()
3763            .unwrap()
3764            .unwrap();
3765
3766        server.abort();
3767    }
3768
3769    #[tokio::test]
3770    async fn post_error_deletes_initialized_connection() {
3771        let delete_count = Arc::new(AtomicUsize::new(0));
3772        let delete_count_for_handler = delete_count.clone();
3773        let app = Router::new().route(
3774            "/acp",
3775            post(initialize_response).get(pending_sse).delete(move || {
3776                let delete_count = delete_count_for_handler.clone();
3777                async move {
3778                    delete_count.fetch_add(1, Ordering::SeqCst);
3779                    StatusCode::ACCEPTED
3780                }
3781            }),
3782        );
3783        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3784        let addr = listener.local_addr().unwrap();
3785        let server = tokio::spawn(async move {
3786            axum::serve(listener, app).await.unwrap();
3787        });
3788        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3789        let (mut caller, transport) = Channel::duplex();
3790        let transport = tokio::spawn(run(client, transport));
3791
3792        caller
3793            .tx
3794            .unbounded_send(single_frame(
3795                RawJsonRpcMessage::request(
3796                    "initialize".to_string(),
3797                    json!({}),
3798                    RequestId::Number(1),
3799                )
3800                .unwrap(),
3801            ))
3802            .unwrap();
3803        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3804            .await
3805            .unwrap()
3806            .unwrap()
3807            .unwrap();
3808        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3809
3810        caller
3811            .tx
3812            .unbounded_send(single_frame(
3813                RawJsonRpcMessage::request(
3814                    "session/prompt".to_string(),
3815                    json!({}),
3816                    RequestId::Number(2),
3817                )
3818                .unwrap(),
3819            ))
3820            .unwrap();
3821        let error = timeout(Duration::from_secs(1), transport)
3822            .await
3823            .unwrap()
3824            .unwrap()
3825            .unwrap_err();
3826
3827        assert!(error.to_string().contains("POST"));
3828        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3829
3830        server.abort();
3831    }
3832
3833    #[tokio::test]
3834    async fn connection_sse_disconnect_fails_transport() {
3835        let delete_count = Arc::new(AtomicUsize::new(0));
3836        let delete_count_for_handler = delete_count.clone();
3837        let app = Router::new().route(
3838            "/acp",
3839            post(initialize_response).get(closed_sse).delete(move || {
3840                let delete_count = delete_count_for_handler.clone();
3841                async move {
3842                    delete_count.fetch_add(1, Ordering::SeqCst);
3843                    StatusCode::ACCEPTED
3844                }
3845            }),
3846        );
3847        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3848        let addr = listener.local_addr().unwrap();
3849        let server = tokio::spawn(async move {
3850            axum::serve(listener, app).await.unwrap();
3851        });
3852        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3853        let (mut caller, transport) = Channel::duplex();
3854        let transport = tokio::spawn(run(client, transport));
3855
3856        caller
3857            .tx
3858            .unbounded_send(single_frame(
3859                RawJsonRpcMessage::request(
3860                    "initialize".to_string(),
3861                    json!({}),
3862                    RequestId::Number(1),
3863                )
3864                .unwrap(),
3865            ))
3866            .unwrap();
3867        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3868            .await
3869            .unwrap()
3870            .unwrap()
3871            .unwrap();
3872        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3873
3874        let error = timeout(Duration::from_secs(1), transport)
3875            .await
3876            .unwrap()
3877            .unwrap()
3878            .unwrap_err();
3879
3880        assert!(error.to_string().contains("SSE"));
3881        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3882
3883        server.abort();
3884    }
3885
3886    #[tokio::test]
3887    async fn malformed_sse_json_is_delivered_and_transport_continues() {
3888        let delete_count = Arc::new(AtomicUsize::new(0));
3889        let delete_count_for_handler = delete_count.clone();
3890        let app = Router::new().route(
3891            "/acp",
3892            post(initialize_response)
3893                .get(malformed_sse)
3894                .delete(move || {
3895                    let delete_count = delete_count_for_handler.clone();
3896                    async move {
3897                        delete_count.fetch_add(1, Ordering::SeqCst);
3898                        StatusCode::ACCEPTED
3899                    }
3900                }),
3901        );
3902        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3903        let addr = listener.local_addr().unwrap();
3904        let server = tokio::spawn(async move {
3905            axum::serve(listener, app).await.unwrap();
3906        });
3907        let client = HttpClient::new(format!("http://{addr}")).unwrap();
3908        let (mut caller, transport) = Channel::duplex();
3909        let transport = tokio::spawn(run(client, transport));
3910
3911        caller
3912            .tx
3913            .unbounded_send(single_frame(
3914                RawJsonRpcMessage::request(
3915                    "initialize".to_string(),
3916                    json!({}),
3917                    RequestId::Number(1),
3918                )
3919                .unwrap(),
3920            ))
3921            .unwrap();
3922        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
3923            .await
3924            .unwrap()
3925            .unwrap()
3926            .unwrap();
3927        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
3928
3929        let frame = timeout(Duration::from_secs(1), caller.rx.next())
3930            .await
3931            .unwrap()
3932            .unwrap();
3933
3934        let TransportFrame::Malformed { raw, error } = frame else {
3935            panic!("expected malformed frame, got {frame:?}");
3936        };
3937        assert_eq!(raw, "{not json");
3938        assert_eq!(error.code, AcpError::parse_error().code);
3939        drop(caller);
3940        timeout(Duration::from_secs(1), transport)
3941            .await
3942            .unwrap()
3943            .unwrap()
3944            .unwrap();
3945        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
3946
3947        server.abort();
3948    }
3949
3950    #[tokio::test]
3951    async fn malformed_ws_json_reports_parse_error_and_continues() {
3952        let app = Router::new().route("/acp", get(malformed_then_valid_ws));
3953        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
3954        let addr = listener.local_addr().unwrap();
3955        let server = tokio::spawn(async move {
3956            axum::serve(listener, app).await.unwrap();
3957        });
3958        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
3959        let (mut caller, transport) = Channel::duplex();
3960        let transport = tokio::spawn(run(client, transport));
3961
3962        let frame = timeout(Duration::from_secs(1), caller.rx.next())
3963            .await
3964            .unwrap()
3965            .unwrap();
3966        let TransportFrame::Malformed { raw, error } = frame else {
3967            panic!("expected malformed frame, got {frame:?}");
3968        };
3969        assert_eq!(raw, "{not json");
3970        assert_eq!(error.code, AcpError::parse_error().code);
3971
3972        let message = timeout(Duration::from_secs(1), caller.rx.next())
3973            .await
3974            .unwrap()
3975            .unwrap()
3976            .unwrap();
3977        assert!(matches!(message, RawJsonRpcMessage::Response(_)));
3978
3979        drop(caller);
3980        timeout(Duration::from_secs(1), transport)
3981            .await
3982            .unwrap()
3983            .unwrap()
3984            .unwrap();
3985
3986        server.abort();
3987    }
3988
3989    fn valid_ws_response_headers() -> HeaderMap {
3990        HeaderMap::from_iter([
3991            (
3992                reqwest::header::UPGRADE,
3993                HeaderValue::from_static("websocket"),
3994            ),
3995            (
3996                reqwest::header::CONNECTION,
3997                HeaderValue::from_static("Upgrade"),
3998            ),
3999            (
4000                reqwest::header::SEC_WEBSOCKET_ACCEPT,
4001                HeaderValue::from_static("s3pPLMBiTxaQ9kYGzzhZRbK+xOo="),
4002            ),
4003        ])
4004    }
4005
4006    #[test]
4007    fn websocket_response_validation() {
4008        let valid = valid_ws_response_headers();
4009        let validate = |version, status, headers: &HeaderMap| {
4010            validate_ws_response(version, status, headers, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=")
4011        };
4012        let version = reqwest::Version::HTTP_11;
4013        let status = StatusCode::SWITCHING_PROTOCOLS;
4014        validate(version, status, &valid).unwrap();
4015        for version in [
4016            reqwest::Version::HTTP_10,
4017            reqwest::Version::HTTP_2,
4018            reqwest::Version::HTTP_3,
4019        ] {
4020            assert!(validate(version, status, &valid).is_err());
4021        }
4022        for status in [StatusCode::OK, StatusCode::BAD_REQUEST, StatusCode::FOUND] {
4023            assert!(validate(version, status, &valid).is_err());
4024        }
4025        for (header, invalid_values) in [
4026            (
4027                "upgrade",
4028                vec!["", "h2c", "websocket/13", "notwebsocket", "websocket, h2c"],
4029            ),
4030            ("connection", vec!["", "keep-alive", "notupgrade"]),
4031            (
4032                "sec-websocket-accept",
4033                vec!["", "wrong", "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=, wrong"],
4034            ),
4035        ] {
4036            let mut headers = valid.clone();
4037            headers.remove(header);
4038            assert!(validate(version, status, &headers).is_err(), "{header}");
4039            for value in invalid_values {
4040                headers.insert(header, HeaderValue::from_str(value).unwrap());
4041                assert!(
4042                    validate(version, status, &headers).is_err(),
4043                    "{header}: {value}"
4044                );
4045            }
4046            headers.insert(header, HeaderValue::from_bytes(b"\xff").unwrap());
4047            assert!(validate(version, status, &headers).is_err(), "{header}");
4048        }
4049        for header in ["upgrade", "sec-websocket-accept"] {
4050            let mut duplicate = valid.clone();
4051            duplicate.append(header, valid[header].clone());
4052            assert!(validate(version, status, &duplicate).is_err(), "{header}");
4053        }
4054
4055        for header in ["sec-websocket-protocol", "sec-websocket-extensions"] {
4056            for value in ["", "acp", "permessage-deflate"] {
4057                let mut headers = valid.clone();
4058                headers.insert(header, HeaderValue::from_str(value).unwrap());
4059                assert!(validate(version, status, &headers).is_err(), "{header}");
4060            }
4061        }
4062
4063        let mut token_lists = valid;
4064        token_lists.insert("upgrade", HeaderValue::from_static("WebSocket"));
4065        token_lists.insert("connection", HeaderValue::from_static("keep-alive"));
4066        token_lists.append("connection", HeaderValue::from_static("other, uPgRaDe\t "));
4067        validate(version, status, &token_lists).unwrap();
4068    }
4069
4070    #[tokio::test]
4071    async fn websocket_public_transport_validates_before_sending_acp() {
4072        use tokio::io::{AsyncReadExt, AsyncWriteExt};
4073
4074        for case in [
4075            "valid",
4076            "version",
4077            "status",
4078            "upgrade",
4079            "connection",
4080            "accept",
4081            "duplicate-accept",
4082            "subprotocol",
4083            "extension",
4084        ] {
4085            let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4086            let addr = listener.local_addr().unwrap();
4087            let fixture = async {
4088                let (mut socket, _) = listener.accept().await.unwrap();
4089                // Read exactly the handshake, preserving any subsequent ACP
4090                // bytes so the assertion also detects premature writes.
4091                let mut request = Vec::new();
4092                while !request.ends_with(b"\r\n\r\n") {
4093                    request.push(socket.read_u8().await.unwrap());
4094                    assert!(request.len() < 16 * 1024);
4095                }
4096                let request = String::from_utf8(request).unwrap();
4097                let key = request
4098                    .lines()
4099                    .filter_map(|line| line.split_once(':'))
4100                    .find(|(name, _)| name.eq_ignore_ascii_case("sec-websocket-key"))
4101                    .unwrap()
4102                    .1
4103                    .trim();
4104                let accept =
4105                    async_tungstenite::tungstenite::handshake::derive_accept_key(key.as_bytes());
4106                let valid = format!(
4107                    "HTTP/1.1 101 Switching Protocols\r\nUpgrade: websocket\r\nConnection: Upgrade\r\nSec-WebSocket-Accept: {accept}\r\n"
4108                );
4109                let response = match case {
4110                    "valid" => valid,
4111                    "version" => valid.replace("HTTP/1.1", "HTTP/1.0"),
4112                    "status" => valid.replace("101 Switching Protocols", "200 OK"),
4113                    "upgrade" => valid.replace("Upgrade: websocket", "Upgrade: not-websocket"),
4114                    "connection" => valid.replace("Connection: Upgrade", "Connection: keep-alive"),
4115                    "accept" => valid.replace(&accept, "wrong"),
4116                    "duplicate-accept" => format!("{valid}Sec-WebSocket-Accept: {accept}\r\n"),
4117                    "subprotocol" => format!("{valid}Sec-WebSocket-Protocol: acp\r\n"),
4118                    "extension" => {
4119                        format!("{valid}Sec-WebSocket-Extensions: permessage-deflate\r\n")
4120                    }
4121                    _ => unreachable!(),
4122                };
4123                socket
4124                    .write_all(format!("{response}\r\n").as_bytes())
4125                    .await
4126                    .unwrap();
4127                let mut received = Vec::new();
4128                socket.read_to_end(&mut received).await.unwrap();
4129                received
4130            };
4131            let client = HttpClient::new(format!("ws://{addr}")).unwrap();
4132            let (caller, transport) = ConnectTo::<Client>::into_channel_and_future(client);
4133            let transport = transport.expect("HttpClient owns its transport driver");
4134            caller
4135                .tx
4136                .unbounded_send(single_frame(
4137                    RawJsonRpcMessage::notification("custom/queued".to_string(), json!({}))
4138                        .unwrap(),
4139                ))
4140                .unwrap();
4141            drop(caller);
4142
4143            // Neither future is spawned: timeout drops both fixtures and
4144            // sockets together, without detached tasks or synchronization sleeps.
4145            let (result, received) = timeout(Duration::from_secs(2), async {
4146                futures::join!(transport, fixture)
4147            })
4148            .await
4149            .expect("handshake fixture should complete");
4150            if case == "valid" {
4151                result.unwrap();
4152                assert!(!received.is_empty(), "valid handshake must send queued ACP");
4153                assert_eq!(received[0], 0x81, "first frame must be WebSocket text");
4154            } else {
4155                assert!(result.is_err(), "{case}: invalid handshake accepted");
4156                assert!(
4157                    received.is_empty(),
4158                    "{case}: ACP escaped before handshake validation"
4159                );
4160            }
4161        }
4162    }
4163
4164    #[tokio::test]
4165    async fn websocket_serializes_batch_as_one_text_frame() {
4166        let (caller, transport) = Channel::duplex();
4167        let Channel {
4168            tx: outgoing,
4169            rx: incoming,
4170        } = caller;
4171        drop(incoming);
4172        outgoing
4173            .unbounded_send(TransportFrame::Batch(
4174                TransportBatch::from_messages([
4175                    RawJsonRpcMessage::notification("custom/first".to_string(), json!({})).unwrap(),
4176                    RawJsonRpcMessage::notification("custom/second".to_string(), json!({}))
4177                        .unwrap(),
4178                ])
4179                .unwrap(),
4180            ))
4181            .unwrap();
4182        drop(outgoing);
4183
4184        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
4185        timeout(
4186            Duration::from_secs(1),
4187            drive_ws(
4188                RecordingWsSink(ws_output_tx),
4189                futures::stream::pending::<Result<WsMessage, std::io::Error>>(),
4190                transport,
4191            ),
4192        )
4193        .await
4194        .unwrap()
4195        .unwrap();
4196        let frames = ws_output.by_ref().collect::<Vec<_>>().await;
4197
4198        let WsMessage::Text(text) = &frames[0] else {
4199            panic!("batch was not sent as WebSocket text");
4200        };
4201        let batch = serde_json::from_str::<serde_json::Value>(text.as_str()).unwrap();
4202        let entries = batch.as_array().expect("batch should remain an array");
4203        assert_eq!(entries.len(), 2);
4204        assert_eq!(entries[0]["method"], "custom/first");
4205        assert_eq!(entries[1]["method"], "custom/second");
4206        assert!(matches!(frames.get(1), Some(WsMessage::Close(None))));
4207        assert_eq!(frames.len(), 2);
4208    }
4209
4210    #[tokio::test]
4211    async fn websocket_drain_discards_incoming_after_receiver_closes() {
4212        let (caller, transport) = Channel::duplex();
4213        let Channel {
4214            tx: outgoing,
4215            rx: incoming,
4216        } = caller;
4217        drop(incoming);
4218
4219        let inbound =
4220            RawJsonRpcMessage::notification("custom/inbound".to_string(), json!({})).unwrap();
4221        let inbound = WsMessage::Text(serde_json::to_string(&inbound).unwrap().into());
4222        let ws_rx = QueueOutgoingThenText {
4223            text: Some(inbound),
4224            outgoing: Some(outgoing),
4225        };
4226        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
4227        timeout(
4228            Duration::from_secs(1),
4229            drive_ws(RecordingWsSink(ws_output_tx), ws_rx, transport),
4230        )
4231        .await
4232        .unwrap()
4233        .unwrap();
4234        let mut frames = Vec::new();
4235        while let Some(frame) = ws_output.next().await {
4236            frames.push(frame);
4237        }
4238
4239        let messages = frames
4240            .iter()
4241            .filter_map(|frame| match frame {
4242                WsMessage::Text(text) => {
4243                    Some(serde_json::from_str::<RawJsonRpcMessage>(text.as_str()).unwrap())
4244                }
4245                _ => None,
4246            })
4247            .collect::<Vec<_>>();
4248        let methods = messages
4249            .iter()
4250            .filter_map(method_for_message)
4251            .collect::<Vec<_>>();
4252        assert_eq!(methods, ["custom/first", "custom/second"]);
4253        assert!(matches!(frames.last(), Some(WsMessage::Close(None))));
4254    }
4255
4256    #[tokio::test]
4257    async fn websocket_reader_runs_while_send_is_backpressured() {
4258        let (caller, transport) = Channel::duplex();
4259        let Channel {
4260            tx: outgoing,
4261            rx: incoming,
4262        } = caller;
4263        drop(incoming);
4264        outgoing
4265            .unbounded_send(single_frame(
4266                RawJsonRpcMessage::notification("custom/queued".to_string(), json!({})).unwrap(),
4267            ))
4268            .unwrap();
4269        drop(outgoing);
4270
4271        let (started_tx, started_rx) = mpsc::unbounded();
4272        let (release_tx, release_rx) = futures::channel::oneshot::channel();
4273        let (ws_output_tx, mut ws_output) = mpsc::unbounded();
4274        let ws_tx = BackpressuredWsSink {
4275            output: ws_output_tx,
4276            started: started_tx,
4277            release: Some(release_rx),
4278        };
4279        let ws_rx = ReleaseBackpressureOnPoll {
4280            started: started_rx,
4281            release: Some(release_tx),
4282        };
4283
4284        timeout(Duration::from_secs(1), drive_ws(ws_tx, ws_rx, transport))
4285            .await
4286            .expect("WebSocket reader was not polled while its writer was backpressured")
4287            .unwrap();
4288        let frames = ws_output.by_ref().collect::<Vec<_>>().await;
4289
4290        let WsMessage::Text(text) = &frames[0] else {
4291            panic!("queued message was not sent as WebSocket text");
4292        };
4293        let message = serde_json::from_str::<RawJsonRpcMessage>(text.as_str()).unwrap();
4294        assert_eq!(method_for_message(&message), Some("custom/queued"));
4295        assert!(matches!(frames.get(1), Some(WsMessage::Close(None))));
4296        assert_eq!(frames.len(), 2);
4297    }
4298
4299    #[tokio::test]
4300    async fn websocket_finish_preserves_close_failure() {
4301        struct FailingCloseSink;
4302        impl WsSink for FailingCloseSink {
4303            fn send(
4304                &mut self,
4305                message: WsMessage,
4306            ) -> impl std::future::Future<Output = Result<(), String>> + Send {
4307                assert!(matches!(message, WsMessage::Close(None)));
4308                futures::future::ready(Err("close failed".to_string()))
4309            }
4310        }
4311
4312        let (caller, transport) = Channel::duplex();
4313        drop(caller.tx);
4314        let error = drive_ws(
4315            FailingCloseSink,
4316            futures::stream::pending::<Result<WsMessage, Infallible>>(),
4317            transport,
4318        )
4319        .await
4320        .unwrap_err();
4321        assert!(error.to_string().contains("ws close: close failed"));
4322    }
4323
4324    #[tokio::test]
4325    async fn peer_ws_close_fails_transport() {
4326        let app = Router::new().route("/acp", get(close_ws));
4327        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4328        let addr = listener.local_addr().unwrap();
4329        let server = tokio::spawn(async move {
4330            axum::serve(listener, app).await.unwrap();
4331        });
4332        let client = HttpClient::new(format!("ws://{addr}")).unwrap();
4333        let (_caller, transport) = Channel::duplex();
4334        let transport = tokio::spawn(run(client, transport));
4335
4336        let error = timeout(Duration::from_secs(1), transport)
4337            .await
4338            .unwrap()
4339            .unwrap()
4340            .unwrap_err();
4341        assert!(error.to_string().contains("WebSocket closed by peer"));
4342
4343        server.abort();
4344    }
4345
4346    #[tokio::test]
4347    async fn websocket_builder_sends_default_headers() {
4348        let (header_tx, mut header_rx) = tokio::sync::mpsc::unbounded_channel();
4349        let app = Router::new().route(
4350            "/acp",
4351            get(move |headers: HeaderMap, ws: WebSocketUpgrade| {
4352                let header_tx = header_tx.clone();
4353                async move {
4354                    header_tx.send(headers).unwrap();
4355                    ws.on_upgrade(|_socket| async {})
4356                }
4357            }),
4358        );
4359        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4360        let addr = listener.local_addr().unwrap();
4361        let server = tokio::spawn(async move {
4362            axum::serve(listener, app).await.unwrap();
4363        });
4364
4365        let mut default_headers = reqwest::header::HeaderMap::new();
4366        default_headers.insert(
4367            reqwest::header::HeaderName::from_static("x-acp-test-client"),
4368            reqwest::header::HeaderValue::from_static("from-reqwest"),
4369        );
4370        for (name, value) in [
4371            ("connection", "close"),
4372            ("upgrade", "h2c"),
4373            ("sec-websocket-version", "12"),
4374            ("sec-websocket-key", "not-a-websocket-key"),
4375        ] {
4376            default_headers.insert(name, HeaderValue::from_static(value));
4377        }
4378        let client = HttpClient::builder(format!("ws://{addr}"))
4379            .configure_http(|http| http.default_headers(default_headers))
4380            .configure_http(reqwest::ClientBuilder::no_proxy)
4381            .build()
4382            .unwrap();
4383        let (_caller, transport) = Channel::duplex();
4384        let transport = tokio::spawn(run(client, transport));
4385
4386        let headers = timeout(Duration::from_secs(1), header_rx.recv())
4387            .await
4388            .expect("WebSocket handshake should reach the server")
4389            .expect("handshake headers were not captured");
4390        assert_eq!(
4391            headers.get("x-acp-test-client").map(HeaderValue::as_bytes),
4392            Some(&b"from-reqwest"[..]),
4393            "default headers must be retained across configure_http calls and sent on the handshake"
4394        );
4395        assert_eq!(headers["connection"], "Upgrade");
4396        assert_eq!(headers["upgrade"], "websocket");
4397        assert_eq!(headers["sec-websocket-version"], "13");
4398        assert_ne!(headers["sec-websocket-key"], "not-a-websocket-key");
4399        for name in [
4400            "connection",
4401            "upgrade",
4402            "sec-websocket-version",
4403            "sec-websocket-key",
4404        ] {
4405            assert_eq!(headers.get_all(name).iter().count(), 1, "{name}");
4406        }
4407
4408        transport.abort();
4409        drop(transport.await);
4410        server.abort();
4411        drop(server.await);
4412    }
4413
4414    async fn assert_websocket_handshake_times_out(
4415        configure: impl FnOnce(reqwest::ClientBuilder) -> reqwest::ClientBuilder,
4416    ) {
4417        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4418        let addr = listener.local_addr().unwrap();
4419        let client = HttpClient::builder(format!("ws://{addr}"))
4420            .configure_http(reqwest::ClientBuilder::no_proxy)
4421            .configure_http(configure)
4422            .build()
4423            .unwrap();
4424        let (_caller, transport) = Channel::duplex();
4425
4426        let error = timeout(Duration::from_secs(1), run(client, transport))
4427            .await
4428            .expect("custom reqwest timeout should fail the WebSocket handshake")
4429            .expect_err("handshake should not succeed while the listener never accepts");
4430        assert!(
4431            error.to_string().contains("WebSocket connect failed"),
4432            "{error}"
4433        );
4434
4435        drop(listener);
4436    }
4437
4438    #[tokio::test]
4439    async fn websocket_builder_honors_request_timeout() {
4440        assert_websocket_handshake_times_out(|http| http.timeout(Duration::from_millis(200))).await;
4441    }
4442
4443    #[tokio::test]
4444    async fn websocket_builder_honors_read_timeout() {
4445        assert_websocket_handshake_times_out(|http| http.read_timeout(Duration::from_millis(200)))
4446            .await;
4447    }
4448
4449    #[tokio::test]
4450    async fn dropped_transport_future_deletes_initialized_connection() {
4451        let delete_count = Arc::new(AtomicUsize::new(0));
4452        let delete_count_for_handler = delete_count.clone();
4453        let app = Router::new().route(
4454            "/acp",
4455            post(initialize_response).get(pending_sse).delete(move || {
4456                let delete_count = delete_count_for_handler.clone();
4457                async move {
4458                    delete_count.fetch_add(1, Ordering::SeqCst);
4459                    StatusCode::ACCEPTED
4460                }
4461            }),
4462        );
4463        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4464        let addr = listener.local_addr().unwrap();
4465        let server = tokio::spawn(async move {
4466            axum::serve(listener, app).await.unwrap();
4467        });
4468        let client = HttpClient::new(format!("http://{addr}")).unwrap();
4469        let (mut caller, transport) = Channel::duplex();
4470        let mut transport = Box::pin(run(client, transport));
4471
4472        caller
4473            .tx
4474            .unbounded_send(single_frame(
4475                RawJsonRpcMessage::request(
4476                    "initialize".to_string(),
4477                    json!({}),
4478                    RequestId::Number(1),
4479                )
4480                .unwrap(),
4481            ))
4482            .unwrap();
4483        let init_response = timeout(Duration::from_secs(1), async {
4484            tokio::select! {
4485                result = &mut transport => {
4486                    panic!("transport ended before initialize response: {result:?}");
4487                }
4488                msg = caller.rx.next() => {
4489                    msg.unwrap().unwrap()
4490                }
4491            }
4492        })
4493        .await
4494        .unwrap();
4495        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
4496
4497        drop(transport);
4498        wait_for_delete(&delete_count).await;
4499
4500        server.abort();
4501    }
4502
4503    #[tokio::test]
4504    async fn dropped_transport_during_close_retries_delete() {
4505        let delete_count = Arc::new(AtomicUsize::new(0));
4506        let delete_count_for_handler = delete_count.clone();
4507        let release_delete = Arc::new(Notify::new());
4508        let release_delete_for_handler = release_delete.clone();
4509        let app = Router::new().route(
4510            "/acp",
4511            post(initialize_response).get(pending_sse).delete(move || {
4512                let delete_count = delete_count_for_handler.clone();
4513                let release_delete = release_delete_for_handler.clone();
4514                async move {
4515                    delete_count.fetch_add(1, Ordering::SeqCst);
4516                    release_delete.notified().await;
4517                    StatusCode::ACCEPTED
4518                }
4519            }),
4520        );
4521        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4522        let addr = listener.local_addr().unwrap();
4523        let server = tokio::spawn(async move {
4524            axum::serve(listener, app).await.unwrap();
4525        });
4526        let client = HttpClient::new(format!("http://{addr}")).unwrap();
4527        let (mut caller, transport) = Channel::duplex();
4528        let transport = tokio::spawn(run(client, transport));
4529
4530        caller
4531            .tx
4532            .unbounded_send(single_frame(
4533                RawJsonRpcMessage::request(
4534                    "initialize".to_string(),
4535                    json!({}),
4536                    RequestId::Number(1),
4537                )
4538                .unwrap(),
4539            ))
4540            .unwrap();
4541        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
4542            .await
4543            .unwrap()
4544            .unwrap()
4545            .unwrap();
4546        assert!(matches!(init_response, RawJsonRpcMessage::Response(_)));
4547
4548        drop(caller);
4549        wait_for_delete_count(&delete_count, 1).await;
4550        transport.abort();
4551        wait_for_delete_count(&delete_count, 2).await;
4552        release_delete.notify_waiters();
4553        drop(transport.await);
4554
4555        server.abort();
4556    }
4557
4558    #[tokio::test]
4559    async fn initialize_error_without_connection_id_is_delivered_without_sse() {
4560        let get_count = Arc::new(AtomicUsize::new(0));
4561        let get_count_for_handler = get_count.clone();
4562        let app = Router::new().route(
4563            "/acp",
4564            post(initialize_error_response).get(move || {
4565                let get_count = get_count_for_handler.clone();
4566                async move {
4567                    get_count.fetch_add(1, Ordering::SeqCst);
4568                    pending_sse().await
4569                }
4570            }),
4571        );
4572        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4573        let addr = listener.local_addr().unwrap();
4574        let server = tokio::spawn(async move {
4575            axum::serve(listener, app).await.unwrap();
4576        });
4577        let client = HttpClient::new(format!("http://{addr}")).unwrap();
4578        let (mut caller, transport) = Channel::duplex();
4579        let transport = tokio::spawn(run(client, transport));
4580
4581        caller
4582            .tx
4583            .unbounded_send(single_frame(
4584                RawJsonRpcMessage::request(
4585                    "initialize".to_string(),
4586                    json!({}),
4587                    RequestId::Number(1),
4588                )
4589                .unwrap(),
4590            ))
4591            .unwrap();
4592        let init_response = timeout(Duration::from_secs(1), caller.rx.next())
4593            .await
4594            .unwrap()
4595            .unwrap()
4596            .unwrap();
4597
4598        assert!(matches!(
4599            init_response,
4600            RawJsonRpcMessage::Response(RpcResponse::Error {
4601                id: RequestId::Number(1),
4602                ..
4603            })
4604        ));
4605        assert_eq!(get_count.load(Ordering::SeqCst), 0);
4606
4607        drop(caller);
4608        timeout(Duration::from_secs(1), transport)
4609            .await
4610            .unwrap()
4611            .unwrap()
4612            .unwrap();
4613
4614        server.abort();
4615    }
4616
4617    #[tokio::test]
4618    async fn malformed_initialize_body_with_connection_id_is_deleted() {
4619        let delete_count = Arc::new(AtomicUsize::new(0));
4620        let delete_count_for_handler = delete_count.clone();
4621        let app = Router::new().route(
4622            "/acp",
4623            post(malformed_initialize_response).delete(move || {
4624                let delete_count = delete_count_for_handler.clone();
4625                async move {
4626                    delete_count.fetch_add(1, Ordering::SeqCst);
4627                    StatusCode::ACCEPTED
4628                }
4629            }),
4630        );
4631        let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
4632        let addr = listener.local_addr().unwrap();
4633        let server = tokio::spawn(async move {
4634            axum::serve(listener, app).await.unwrap();
4635        });
4636        let client = HttpClient::new(format!("http://{addr}")).unwrap();
4637        let (caller, transport) = Channel::duplex();
4638        let transport = tokio::spawn(run(client, transport));
4639
4640        caller
4641            .tx
4642            .unbounded_send(single_frame(
4643                RawJsonRpcMessage::request(
4644                    "initialize".to_string(),
4645                    json!({}),
4646                    RequestId::Number(1),
4647                )
4648                .unwrap(),
4649            ))
4650            .unwrap();
4651        let error = timeout(Duration::from_secs(1), transport)
4652            .await
4653            .unwrap()
4654            .unwrap()
4655            .unwrap_err();
4656
4657        assert!(error.to_string().contains("initialize"));
4658        wait_for_delete(&delete_count).await;
4659
4660        server.abort();
4661    }
4662
4663    async fn wait_for_delete(delete_count: &AtomicUsize) {
4664        wait_for_delete_count(delete_count, 1).await;
4665        assert_eq!(delete_count.load(Ordering::SeqCst), 1);
4666    }
4667
4668    async fn wait_for_delete_count(delete_count: &AtomicUsize, expected: usize) {
4669        timeout(Duration::from_secs(1), async {
4670            loop {
4671                if delete_count.load(Ordering::SeqCst) >= expected {
4672                    break;
4673                }
4674                sleep(Duration::from_millis(10)).await;
4675            }
4676        })
4677        .await
4678        .unwrap();
4679    }
4680
4681    async fn initialize_response() -> impl IntoResponse {
4682        let mut headers = HeaderMap::new();
4683        headers.insert(HEADER_CONNECTION_ID, HeaderValue::from_static("conn-1"));
4684        (
4685            StatusCode::OK,
4686            headers,
4687            Json(RawJsonRpcMessage::response(
4688                RequestId::Number(1),
4689                Ok(json!({})),
4690            )),
4691        )
4692    }
4693
4694    async fn initialize_error_response() -> Json<RawJsonRpcMessage> {
4695        Json(RawJsonRpcMessage::response(
4696            RequestId::Number(1),
4697            Err(AcpError::invalid_request().data("initialize rejected")),
4698        ))
4699    }
4700
4701    async fn malformed_initialize_response() -> impl IntoResponse {
4702        let mut headers = HeaderMap::new();
4703        headers.insert(HEADER_CONNECTION_ID, HeaderValue::from_static("conn-1"));
4704        (StatusCode::OK, headers, "{not json")
4705    }
4706
4707    async fn pending_sse() -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
4708        Sse::new(futures::stream::pending())
4709    }
4710
4711    fn sse_event(message: RawJsonRpcMessage) -> Event {
4712        Event::default().data(serde_json::to_string(&message).unwrap())
4713    }
4714
4715    async fn malformed_sse() -> Sse<impl futures::Stream<Item = Result<Event, Infallible>>> {
4716        let invalid = futures::stream::once(async {
4717            Ok::<_, Infallible>(Event::default().data("{not json"))
4718        });
4719        Sse::new(invalid.chain(futures::stream::pending()))
4720    }
4721
4722    async fn malformed_then_valid_ws(ws: WebSocketUpgrade) -> impl IntoResponse {
4723        ws.on_upgrade(|mut socket| async move {
4724            drop(socket.send(AxumWsMessage::Text("{not json".into())).await);
4725            let valid = serde_json::to_string(&RawJsonRpcMessage::response(
4726                RequestId::Number(1),
4727                Ok(json!({})),
4728            ))
4729            .unwrap();
4730            drop(socket.send(AxumWsMessage::Text(valid.into())).await);
4731            futures::future::pending::<()>().await;
4732        })
4733    }
4734
4735    async fn close_ws(ws: WebSocketUpgrade) -> impl IntoResponse {
4736        ws.on_upgrade(|mut socket| async move {
4737            drop(socket.send(AxumWsMessage::Close(None)).await);
4738        })
4739    }
4740
4741    async fn closed_sse() -> StatusCode {
4742        StatusCode::OK
4743    }
4744}