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