use std::io::{Error, ErrorKind, Result};
use bytes::Bytes;
use futures_util::{
SinkExt, StreamExt,
stream::{SplitSink, SplitStream},
};
use tokio::net::{TcpListener, TcpStream};
use tokio_tungstenite::{
MaybeTlsStream, WebSocketStream, accept_async, connect_async,
tungstenite::{Message as WebSocketMessage, error::Error as WebSocketError},
};
use tracing::info;
use super::{IoFuture, IpcConnection, IpcListener};
#[derive(Debug)]
pub struct WebSocketListener {
listener: TcpListener,
local_addr: String,
}
impl WebSocketListener {
pub async fn bind(local_addr: &str) -> Result<Self> {
let listener = TcpListener::bind(local_addr).await?;
Ok(Self {
listener,
local_addr: local_addr.to_string(),
})
}
}
impl IpcListener for WebSocketListener {
fn local_endpoint(&self) -> &str {
self.local_addr.as_str()
}
fn accept(&self) -> IoFuture<'_, Box<dyn IpcConnection>> {
Box::pin(async move {
let (socket, peer_addr) = self.listener.accept().await?;
let ws_stream =
accept_async(MaybeTlsStream::Plain(socket))
.await
.map_err(|e| match e {
WebSocketError::Io(e) => e,
e => Error::other(e),
})?;
info!("Accepted a new websocket connection from {}", peer_addr);
Ok(
Box::new(WebSocketConnection::new(ws_stream, peer_addr.to_string()))
as Box<dyn IpcConnection>,
)
})
}
}
#[derive(Debug)]
pub struct WebSocketConnection {
peer_addr: String,
tx: SplitSink<WebSocketStream<MaybeTlsStream<TcpStream>>, WebSocketMessage>,
rx: SplitStream<WebSocketStream<MaybeTlsStream<TcpStream>>>,
pending_pong: Option<Bytes>,
}
impl WebSocketConnection {
fn new(ws_stream: WebSocketStream<MaybeTlsStream<TcpStream>>, peer_addr: String) -> Self {
let (tx, rx) = ws_stream.split();
Self {
peer_addr,
tx,
rx,
pending_pong: None,
}
}
async fn send(&mut self, msg: WebSocketMessage) -> Result<()> {
self.tx.send(msg).await.map_err(|e| match e {
WebSocketError::Io(e) => e,
e => Error::other(e),
})
}
async fn recv(&mut self) -> Result<WebSocketMessage> {
self.rx
.next()
.await
.ok_or(ErrorKind::ConnectionAborted)?
.map_err(|e| match e {
WebSocketError::Io(e) => e,
e => Error::other(e),
})
}
}
impl IpcConnection for WebSocketConnection {
async fn connect(peer_addr: &str) -> Result<Self> {
let (ws_stream, _) = connect_async(peer_addr).await.map_err(|e| match e {
WebSocketError::Io(e) => e,
e => Error::other(e),
})?;
info!("Connected to websocket server {}", peer_addr);
Ok(Self::new(ws_stream, peer_addr.to_string()))
}
fn peer_endpoint(&self) -> &str {
self.peer_addr.as_str()
}
fn close(&mut self) -> IoFuture<'_, ()> {
Box::pin(async move {
self.tx.close().await.map_err(|e| match e {
WebSocketError::Io(e) => e,
e => Error::other(e),
})
})
}
fn send(&mut self, buf: Bytes) -> IoFuture<'_, ()> {
Box::pin(self.send(WebSocketMessage::Binary(buf)))
}
fn recv(&mut self) -> IoFuture<'_, Bytes> {
Box::pin(async move {
loop {
if let Some(payload) = self.pending_pong.clone() {
self.send(WebSocketMessage::Pong(payload)).await?; self.pending_pong = None;
}
let message = self.recv().await?;
match message {
WebSocketMessage::Binary(payload) => return Ok(payload),
WebSocketMessage::Ping(payload) => {
self.pending_pong = Some(payload);
}
WebSocketMessage::Pong(_) => {}
WebSocketMessage::Close(_) => {
return Err(Error::new(
ErrorKind::ConnectionAborted,
"received close message",
));
}
_ => return Err(Error::other("received non-binary message")),
}
}
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn test_listener() {
let listener = WebSocketListener::bind("127.0.0.1:0").await.unwrap();
assert_eq!(listener.local_endpoint(), "127.0.0.1:0");
drop(listener);
WebSocketListener::bind("not-a-valid-address")
.await
.expect_err("bind to an invalid address must fail");
let probe = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = probe.local_addr().unwrap();
drop(probe);
WebSocketConnection::connect(&format!("ws://{addr}"))
.await
.expect_err("connect with no server must fail");
}
#[tokio::test]
async fn test_connection() {
let listener = WebSocketListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.listener.local_addr().unwrap();
let url = format!("ws://{addr}");
let server = tokio::spawn(async move {
let mut conn = listener.accept().await.expect("accept first");
assert!(!conn.peer_endpoint().is_empty());
let msg = conn.recv().await.expect("recv");
conn.send(msg).await.expect("send");
let mut conn = listener.accept().await.expect("accept second");
conn.close().await.expect("close");
});
let mut client = WebSocketConnection::connect(&url).await.expect("connect");
assert_eq!(client.peer_endpoint(), url);
let payload = Bytes::from_static(b"hello");
<WebSocketConnection as IpcConnection>::send(&mut client, payload.clone())
.await
.expect("send");
let received = <WebSocketConnection as IpcConnection>::recv(&mut client)
.await
.expect("recv");
assert_eq!(received, payload);
<WebSocketConnection as IpcConnection>::close(&mut client)
.await
.expect("close");
let mut client = WebSocketConnection::connect(&url).await.expect("connect");
let err = <WebSocketConnection as IpcConnection>::recv(&mut client)
.await
.expect_err("recv after server close must fail");
assert_eq!(err.kind(), ErrorKind::ConnectionAborted);
server.await.unwrap();
}
#[tokio::test]
async fn test_websocket_control_frames() {
let tcp = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = tcp.local_addr().unwrap();
let url = format!("ws://{addr}");
let server = tokio::spawn(async move {
let (socket, peer) = tcp.accept().await.unwrap();
let ws = accept_async(MaybeTlsStream::Plain(socket)).await.unwrap();
let mut conn = WebSocketConnection::new(ws, peer.to_string());
conn.send(WebSocketMessage::Ping(Bytes::from_static(b"p")))
.await
.unwrap();
conn.send(WebSocketMessage::Binary(Bytes::from_static(b"hi")))
.await
.unwrap();
let reply = conn.recv().await.unwrap();
assert!(matches!(&reply, WebSocketMessage::Pong(p) if p.as_ref() == b"p"));
let (socket, peer) = tcp.accept().await.unwrap();
let ws = accept_async(MaybeTlsStream::Plain(socket)).await.unwrap();
let mut conn = WebSocketConnection::new(ws, peer.to_string());
conn.send(WebSocketMessage::Pong(Bytes::from_static(b"p")))
.await
.unwrap();
conn.send(WebSocketMessage::Binary(Bytes::from_static(b"ok")))
.await
.unwrap();
let (socket, peer) = tcp.accept().await.unwrap();
let ws = accept_async(MaybeTlsStream::Plain(socket)).await.unwrap();
let mut conn = WebSocketConnection::new(ws, peer.to_string());
conn.send(WebSocketMessage::Text("unexpected".into()))
.await
.unwrap();
});
let mut client = WebSocketConnection::connect(&url).await.unwrap();
let bytes = <WebSocketConnection as IpcConnection>::recv(&mut client)
.await
.unwrap();
assert_eq!(bytes.as_ref(), b"hi");
let mut client = WebSocketConnection::connect(&url).await.unwrap();
let bytes = <WebSocketConnection as IpcConnection>::recv(&mut client)
.await
.unwrap();
assert_eq!(bytes.as_ref(), b"ok");
let mut client = WebSocketConnection::connect(&url).await.unwrap();
let err = <WebSocketConnection as IpcConnection>::recv(&mut client)
.await
.expect_err("text message should be rejected");
assert_ne!(err.kind(), ErrorKind::ConnectionAborted);
server.await.unwrap();
}
}