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