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