Skip to main content

acp_utils/
websocket.rs

1use agent_client_protocol::{ConnectTo, Error, Lines, Role};
2use futures::{SinkExt, StreamExt, stream};
3use std::{io, time::Duration};
4use thiserror::Error;
5use tokio::io::{AsyncRead, AsyncWrite};
6use tokio::sync::mpsc;
7use tokio::time::{Instant, MissedTickBehavior, interval_at, timeout};
8use tokio_tungstenite::WebSocketStream;
9use tokio_tungstenite::tungstenite::protocol::{CloseFrame, frame::coding::CloseCode};
10use tokio_tungstenite::tungstenite::{self, Bytes, Message};
11use tokio_util::sync::PollSender;
12
13#[derive(Debug, Error)]
14pub enum WebSocketError {
15    #[error(transparent)]
16    WebSocket(#[from] tungstenite::Error),
17    #[error("unexpected raw WebSocket frame")]
18    UnexpectedFrame,
19    #[error("binary WebSocket messages are not supported")]
20    BinaryMessage,
21    #[error("WebSocket consumer closed")]
22    ConsumerClosed,
23    #[error("WebSocket write deadline exceeded")]
24    WriteDeadline,
25}
26
27/// WebSocket transport for ACP
28pub struct WebSocketTransport<T> {
29    socket: WebSocketStream<T>,
30}
31
32impl<T: AsyncRead + AsyncWrite + Unpin> WebSocketTransport<T> {
33    pub fn new(socket: WebSocketStream<T>) -> Self {
34        Self { socket }
35    }
36}
37
38impl<T, U> ConnectTo<U> for WebSocketTransport<T>
39where
40    T: AsyncRead + AsyncWrite + Unpin + Send + 'static,
41    U: Role,
42{
43    async fn connect_to(self, peer: impl ConnectTo<U::Counterpart>) -> Result<(), Error> {
44        let (to_tx, to_rx) = mpsc::channel::<String>(32);
45        let (from_tx, from_rx) = mpsc::channel::<io::Result<String>>(32);
46        let output = PollSender::new(to_tx).sink_map_err(|_| io::Error::from(io::ErrorKind::BrokenPipe));
47        let input = stream::unfold(from_rx, |mut rx| async { rx.recv().await.map(|item| (item, rx)) });
48        let connection_future = ConnectTo::<U>::connect_to(Lines::new(output, input), peer);
49        let socket_loop = self.run_socket_loop(to_rx, from_tx);
50        tokio::pin!(connection_future, socket_loop);
51        tokio::select! {
52            result = &mut connection_future => {
53                result?;
54                socket_loop.await.map_err(Error::into_internal_error)
55            }
56
57            result = &mut socket_loop => {
58                result.map_err(Error::into_internal_error)?;
59                connection_future.await
60            },
61        }
62    }
63}
64
65const WRITE_DEADLINE: Duration = Duration::from_secs(10);
66const KEEPALIVE_INTERVAL: Duration = Duration::from_secs(20);
67
68impl<T: AsyncRead + AsyncWrite + Unpin> WebSocketTransport<T> {
69    async fn run_socket_loop(
70        mut self,
71        mut to_rx: mpsc::Receiver<String>,
72        from_tx: mpsc::Sender<io::Result<String>>,
73    ) -> Result<(), WebSocketError> {
74        let mut keepalive = interval_at(Instant::now() + KEEPALIVE_INTERVAL, KEEPALIVE_INTERVAL);
75        keepalive.set_missed_tick_behavior(MissedTickBehavior::Skip);
76        loop {
77            tokio::select! {
78                _ = keepalive.tick() => {
79                    self.write(Message::Ping(Bytes::new())).await?;
80                }
81                message = to_rx.recv() => {
82                    let Some(text) = message else {
83                        return self.finish_close().await;
84                    };
85                    self.write(Message::Text(text.into())).await?;
86                }
87                message = self.socket.next() => {
88                    match message {
89                        Some(Ok(Message::Text(text))) => {
90                            from_tx.send(Ok(text.to_string())).await.map_err(|_| WebSocketError::ConsumerClosed)?;
91                        }
92                        Some(Ok(Message::Ping(_))) => {
93                            // Tungstenite queues the matching Pong; flush it even while ACP is idle.
94                            timeout(WRITE_DEADLINE, self.socket.flush()).await
95                                .map_err(|_| WebSocketError::WriteDeadline)??;
96                        }
97                        Some(Ok(Message::Pong(_))) => {},
98                        Some(Ok(Message::Close(_))) => {
99                            let _ = timeout(WRITE_DEADLINE, self.socket.flush()).await;
100                            return Ok(());
101                        }
102                        Some(Ok(Message::Binary(_))) => {
103                            self.close(CloseCode::Unsupported).await;
104                            return Err(WebSocketError::BinaryMessage);
105                        }
106                        Some(Ok(Message::Frame(_))) => return Err(WebSocketError::UnexpectedFrame),
107                        Some(Err(tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed)) | None => return Ok(()),
108                        Some(Err(error)) => {
109                            if matches!(error, tungstenite::Error::Capacity(_)) {
110                                self.close(CloseCode::Size).await;
111                            }
112                            return Err(error.into());
113                        }
114                    }
115                }
116            }
117        }
118    }
119
120    async fn finish_close(&mut self) -> Result<(), WebSocketError> {
121        timeout(WRITE_DEADLINE, async {
122            match self.socket.send(Message::Close(None)).await {
123                Ok(()) => {}
124                Err(tungstenite::Error::ConnectionClosed | tungstenite::Error::AlreadyClosed) => return Ok(()),
125                Err(error) => return Err(error.into()),
126            }
127            while let Some(message) = self.socket.next().await {
128                match message {
129                    Ok(Message::Close(_)) | Err(tungstenite::Error::ConnectionClosed) => break,
130                    Err(error) => return Err(error.into()),
131                    _ => {}
132                }
133            }
134            Ok(())
135        })
136        .await
137        .map_err(|_| WebSocketError::WriteDeadline)?
138    }
139
140    async fn write(&mut self, message: Message) -> Result<(), WebSocketError> {
141        timeout(WRITE_DEADLINE, self.socket.send(message))
142            .await
143            .map_err(|_| WebSocketError::WriteDeadline)?
144            .map_err(WebSocketError::from)
145    }
146
147    async fn close(&mut self, code: CloseCode) {
148        let _ = self.write(Message::Close(Some(CloseFrame { code, reason: "".into() }))).await;
149    }
150}