acktor_ipc/ipc_method/
websocket.rs1use 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#[derive(Debug)]
28pub struct WebSocketListener {
29 listener: TcpListener,
30 local_addr: String,
31}
32
33impl WebSocketListener {
34 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#[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 if let Some(payload) = self.pending_pong.clone() {
132 self.send(WebSocketMessage::Pong(payload)).await?; 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 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 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 WebSocketListener::bind("not-a-valid-address")
176 .await
177 .expect_err("bind to an invalid address must fail");
178
179 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 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 let mut conn = listener.accept().await?;
205 conn.close().await?;
206
207 Ok::<_, anyhow::Error>(())
208 });
209
210 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 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 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 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 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 let mut client = WebSocketConnection::connect(&url).await?;
270 let bytes = <WebSocketConnection as IpcConnection>::recv(&mut client).await?;
272 assert_eq!(bytes.as_ref(), b"hi");
273
274 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 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}