acktor_ipc/ipc_method/
websocket.rs1use 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#[derive(Debug)]
21pub struct WebSocketListener {
22 listener: TcpListener,
23 local_addr: String,
24}
25
26impl WebSocketListener {
27 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#[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 if let Some(payload) = self.pending_pong.clone() {
141 self.send(WebSocketMessage::Pong(payload)).await?; 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 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 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 WebSocketListener::bind("not-a-valid-address")
182 .await
183 .expect_err("bind to an invalid address must fail");
184
185 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 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 let mut conn = listener.accept().await.expect("accept second");
209 conn.close().await.expect("close");
210 });
211
212 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 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 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 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 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 let mut client = WebSocketConnection::connect(&url).await.unwrap();
279 let bytes = <WebSocketConnection as IpcConnection>::recv(&mut client)
281 .await
282 .unwrap();
283 assert_eq!(bytes.as_ref(), b"hi");
284
285 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 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}