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