Skip to main content

acktor_ipc/ipc_method/
websocket.rs

1//! IPC method implementation using WebSocket.
2
3use std::io::{Error, ErrorKind, Result};
4
5use bytes::Bytes;
6use futures_util::{
7    FutureExt, SinkExt, StreamExt, TryFutureExt,
8    stream::{SplitSink, SplitStream},
9};
10use tokio::net::{TcpListener, TcpStream};
11use tokio_tungstenite::{
12    MaybeTlsStream, WebSocketStream, accept_async, connect_async,
13    tungstenite::{Message as WebSocketMessage, error::Error as WebSocketError},
14};
15use tracing::info;
16
17use super::{IoFuture, IpcConnection, IpcListener};
18
19fn ws_error_to_io_error(e: WebSocketError) -> Error {
20    match e {
21        WebSocketError::Io(e) => e,
22        e => Error::other(e),
23    }
24}
25
26/// IPC listener implemented with WebSocket.
27#[derive(Debug)]
28pub struct WebSocketListener {
29    listener: TcpListener,
30    local_addr: String,
31}
32
33impl WebSocketListener {
34    /// Constructs a new [`WebSocketListener`] with the given bind address.
35    pub async fn bind(local_addr: &str) -> Result<Self> {
36        let listener = TcpListener::bind(local_addr).await?;
37
38        Ok(Self {
39            listener,
40            local_addr: local_addr.to_string(),
41        })
42    }
43}
44
45impl IpcListener for WebSocketListener {
46    fn local_endpoint(&self) -> &str {
47        self.local_addr.as_str()
48    }
49
50    fn accept(&self) -> IoFuture<'_, Box<dyn IpcConnection>> {
51        Box::pin(async move {
52            let (socket, peer_addr) = self.listener.accept().await?;
53
54            let ws_stream = accept_async(MaybeTlsStream::Plain(socket))
55                .await
56                .map_err(ws_error_to_io_error)?;
57
58            info!("Accepted a new websocket connection from {}", peer_addr);
59
60            Ok(
61                Box::new(WebSocketConnection::new(ws_stream, peer_addr.to_string()))
62                    as Box<dyn IpcConnection>,
63            )
64        })
65    }
66}
67
68/// IPC connection implemented with WebSocket.
69#[derive(Debug)]
70pub struct WebSocketConnection {
71    peer_addr: String,
72    tx: SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, WebSocketMessage>,
73    rx: SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>,
74    pending_pong: Option<Bytes>,
75}
76
77impl WebSocketConnection {
78    fn new(ws_stream: WebSocketStream<MaybeTlsStream<TcpStream>>, peer_addr: String) -> Self {
79        let (tx, rx) = ws_stream.split();
80
81        Self {
82            peer_addr,
83            tx,
84            rx,
85            pending_pong: None,
86        }
87    }
88
89    async fn send(&mut self, msg: WebSocketMessage) -> Result<()> {
90        self.tx.send(msg).await.map_err(ws_error_to_io_error)
91    }
92
93    async fn recv(&mut self) -> Result<WebSocketMessage> {
94        self.rx
95            .next()
96            .await
97            .ok_or(ErrorKind::ConnectionAborted)?
98            .map_err(ws_error_to_io_error)
99    }
100}
101
102impl IpcConnection for WebSocketConnection {
103    async fn connect(peer_addr: &str) -> Result<Self> {
104        let (ws_stream, _) = connect_async(peer_addr)
105            .await
106            .map_err(ws_error_to_io_error)?;
107
108        info!("Connected to websocket server {}", peer_addr);
109
110        Ok(Self::new(ws_stream, peer_addr.to_string()))
111    }
112
113    fn peer_endpoint(&self) -> &str {
114        self.peer_addr.as_str()
115    }
116
117    fn close(&mut self) -> IoFuture<'_, ()> {
118        self.tx.close().map_err(ws_error_to_io_error).boxed()
119    }
120
121    fn send(&mut self, buf: Bytes) -> IoFuture<'_, ()> {
122        self.send(WebSocketMessage::Binary(buf)).boxed()
123    }
124
125    fn recv(&mut self) -> IoFuture<'_, Bytes> {
126        async move {
127            loop {
128                // send any buffered Pong, the payload is cloned so `self.pending_pong` keeps
129                // the value until the send completes
130                // NOTE: clone a Bytes is cheap
131                if let Some(payload) = self.pending_pong.clone() {
132                    self.send(WebSocketMessage::Pong(payload)).await?; // #1
133                    self.pending_pong = None;
134                }
135
136                let message = self.recv().await?;
137
138                match message {
139                    WebSocketMessage::Binary(payload) => return Ok(payload),
140                    WebSocketMessage::Ping(payload) => {
141                        // buffer the payload for sending in the next loop
142                        // if the call is cancelled at #1, the payload is still buffered
143                        // and next call retries
144                        self.pending_pong = Some(payload);
145                    }
146                    WebSocketMessage::Pong(_) => {}
147                    WebSocketMessage::Close(_) => {
148                        return Err(Error::new(
149                            ErrorKind::ConnectionAborted,
150                            "received close message",
151                        ));
152                    }
153                    _ => return Err(Error::other("received non-binary message")),
154                }
155            }
156        }
157        .boxed()
158    }
159}
160
161#[cfg(test)]
162mod tests {
163    use pretty_assertions::{assert_eq, assert_ne};
164
165    use super::*;
166
167    #[tokio::test]
168    async fn test_listener() -> anyhow::Result<()> {
169        // bind success + local_endpoint reflects the string passed to `bind`
170        let listener = WebSocketListener::bind("127.0.0.1:0").await?;
171        assert_eq!(listener.local_endpoint(), "127.0.0.1:0");
172        drop(listener);
173
174        // bind to an invalid address fails
175        WebSocketListener::bind("not-a-valid-address")
176            .await
177            .expect_err("bind to an invalid address must fail");
178
179        // connect with no server on the port fails
180        let probe = TcpListener::bind("127.0.0.1:0").await?;
181        let addr = probe.local_addr()?;
182        drop(probe);
183        WebSocketConnection::connect(&format!("ws://{}", addr))
184            .await
185            .expect_err("connect with no server must fail");
186
187        Ok(())
188    }
189
190    #[tokio::test]
191    async fn test_connection() -> anyhow::Result<()> {
192        let listener = WebSocketListener::bind("127.0.0.1:0").await?;
193        let addr = listener.listener.local_addr()?;
194        let url = format!("ws://{}", addr);
195
196        let server = tokio::spawn(async move {
197            // 1st client: echo roundtrip
198            let mut conn = listener.accept().await?;
199            assert!(!conn.peer_endpoint().is_empty());
200            let msg = conn.recv().await?;
201            conn.send(msg).await?;
202
203            // 2nd client: server closes the connection
204            let mut conn = listener.accept().await?;
205            conn.close().await?;
206
207            Ok::<_, anyhow::Error>(())
208        });
209
210        // client 1: full roundtrip, then IpcConnection::close
211        let mut client = WebSocketConnection::connect(&url).await?;
212        assert_eq!(client.peer_endpoint(), url);
213        let payload = Bytes::from_static(b"hello");
214        <WebSocketConnection as IpcConnection>::send(&mut client, payload.clone()).await?;
215        let received = <WebSocketConnection as IpcConnection>::recv(&mut client).await?;
216        assert_eq!(received, payload);
217        <WebSocketConnection as IpcConnection>::close(&mut client).await?;
218
219        // client 2: server closes, recv returns ConnectionAborted
220        let mut client = WebSocketConnection::connect(&url).await?;
221        let err = <WebSocketConnection as IpcConnection>::recv(&mut client)
222            .await
223            .expect_err("recv after server close must fail");
224        assert_eq!(err.kind(), ErrorKind::ConnectionAborted);
225
226        server.await??;
227
228        Ok(())
229    }
230
231    #[tokio::test]
232    async fn test_websocket_control_frames() -> anyhow::Result<()> {
233        let tcp = TcpListener::bind("127.0.0.1:0").await?;
234        let addr = tcp.local_addr()?;
235        let url = format!("ws://{}", addr);
236
237        let server = tokio::spawn(async move {
238            // session 1: Ping -> client should reply with Pong, then consume the Binary
239            let (socket, peer) = tcp.accept().await?;
240            let ws = accept_async(MaybeTlsStream::Plain(socket)).await?;
241            let mut conn = WebSocketConnection::new(ws, peer.to_string());
242            conn.send(WebSocketMessage::Ping(Bytes::from_static(b"p")))
243                .await?;
244            conn.send(WebSocketMessage::Binary(Bytes::from_static(b"hi")))
245                .await?;
246            let reply = conn.recv().await?;
247            assert!(matches!(&reply, WebSocketMessage::Pong(p) if p.as_ref() == b"p"));
248
249            // session 2: Pong is ignored, next Binary returned
250            let (socket, peer) = tcp.accept().await?;
251            let ws = accept_async(MaybeTlsStream::Plain(socket)).await?;
252            let mut conn = WebSocketConnection::new(ws, peer.to_string());
253            conn.send(WebSocketMessage::Pong(Bytes::from_static(b"p")))
254                .await?;
255            conn.send(WebSocketMessage::Binary(Bytes::from_static(b"ok")))
256                .await?;
257
258            // session 3: Text frame is rejected (not ConnectionAborted)
259            let (socket, peer) = tcp.accept().await?;
260            let ws = accept_async(MaybeTlsStream::Plain(socket)).await?;
261            let mut conn = WebSocketConnection::new(ws, peer.to_string());
262            conn.send(WebSocketMessage::Text("unexpected".into()))
263                .await?;
264
265            Ok::<_, anyhow::Error>(())
266        });
267
268        // session 1
269        let mut client = WebSocketConnection::connect(&url).await?;
270        // pong is automatically handled
271        let bytes = <WebSocketConnection as IpcConnection>::recv(&mut client).await?;
272        assert_eq!(bytes.as_ref(), b"hi");
273
274        // session 2
275        let mut client = WebSocketConnection::connect(&url).await?;
276        let bytes = <WebSocketConnection as IpcConnection>::recv(&mut client).await?;
277        assert_eq!(bytes.as_ref(), b"ok");
278
279        // session 3
280        let mut client = WebSocketConnection::connect(&url).await?;
281        let err = <WebSocketConnection as IpcConnection>::recv(&mut client)
282            .await
283            .expect_err("text message should be rejected");
284        assert_ne!(err.kind(), ErrorKind::ConnectionAborted);
285
286        server.await??;
287
288        Ok(())
289    }
290}