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