use std::sync::Arc;
use std::time::Duration;
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
use axum::http::HeaderMap;
use axum::response::Response;
use axum::routing::{any, MethodRouter};
use futures::stream::{SplitSink, SplitStream, StreamExt};
use futures::SinkExt;
use tokio::sync::Mutex;
type WsSink = SplitSink<WebSocket, Message>;
type WsStream = SplitStream<WebSocket>;
pub const DEFAULT_IDLE_PING_MS: u64 = 30_000;
pub const DEFAULT_MAX_MESSAGE_BYTES: usize = 1024 * 1024;
#[derive(Clone, Debug)]
pub struct WsConfig {
pub subprotocols: Vec<String>,
pub idle_ping: Option<Duration>,
pub max_message_bytes: usize,
}
impl Default for WsConfig {
fn default() -> Self {
Self {
subprotocols: Vec::new(),
idle_ping: Some(Duration::from_millis(DEFAULT_IDLE_PING_MS)),
max_message_bytes: DEFAULT_MAX_MESSAGE_BYTES,
}
}
}
pub struct WsSession {
sink: Arc<Mutex<WsSink>>,
stream: Mutex<WsStream>,
pub negotiated_subprotocol: Option<String>,
pub max_message_bytes: usize,
}
impl WsSession {
pub async fn recv(&self) -> Result<Option<WsMessage>, WsError> {
loop {
let next = self.stream.lock().await.next().await;
match next {
Some(Ok(Message::Text(text))) => {
if text.len() > self.max_message_bytes {
self.send_close(1009, "message too big").await?;
return Err(WsError::MessageTooBig);
}
return Ok(Some(WsMessage::Text(text.to_string())));
}
Some(Ok(Message::Binary(bytes))) => {
if bytes.len() > self.max_message_bytes {
self.send_close(1009, "message too big").await?;
return Err(WsError::MessageTooBig);
}
return Ok(Some(WsMessage::Binary(bytes.to_vec())));
}
Some(Ok(Message::Ping(_) | Message::Pong(_))) => continue,
Some(Ok(Message::Close(_))) | None => return Ok(None),
Some(Err(error)) => return Err(WsError::Transport(error.to_string())),
}
}
}
pub async fn send(&self, text: impl Into<String>) -> Result<(), WsError> {
let text = text.into();
if text.len() > self.max_message_bytes {
return Err(WsError::MessageTooBig);
}
send_message(&self.sink, Message::Text(text.into())).await
}
pub async fn send_binary(&self, bytes: Vec<u8>) -> Result<(), WsError> {
if bytes.len() > self.max_message_bytes {
return Err(WsError::MessageTooBig);
}
send_message(&self.sink, Message::Binary(bytes.into())).await
}
pub async fn ping(&self) -> Result<(), WsError> {
send_message(&self.sink, Message::Ping(Default::default())).await
}
pub async fn close(&self, code: u16, reason: &str) -> Result<(), WsError> {
self.send_close(code, reason).await
}
async fn send_close(&self, code: u16, reason: &str) -> Result<(), WsError> {
send_message(
&self.sink,
Message::Close(Some(axum::extract::ws::CloseFrame {
code,
reason: reason.to_string().into(),
})),
)
.await
}
}
async fn send_message(sink: &Mutex<WsSink>, message: Message) -> Result<(), WsError> {
sink.lock()
.await
.send(message)
.await
.map_err(|error| WsError::Transport(error.to_string()))
}
#[derive(Debug, Clone)]
pub enum WsMessage {
Text(String),
Binary(Vec<u8>),
}
#[derive(Debug)]
pub enum WsError {
Transport(String),
MessageTooBig,
}
impl std::fmt::Display for WsError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Transport(message) => write!(f, "ws transport error: {message}"),
Self::MessageTooBig => write!(f, "ws message too big"),
}
}
}
impl std::error::Error for WsError {}
pub fn ws_route<S, F, Fut>(handler: F, config: WsConfig) -> MethodRouter<S>
where
S: Clone + Send + Sync + 'static,
F: Fn(WsSession) -> Fut + Clone + Send + Sync + 'static,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let config = Arc::new(config);
let handler: Arc<dyn Fn(WsSession) -> WsHandlerFuture + Send + Sync + 'static> =
Arc::new(move |session| Box::pin(handler(session)) as _);
any(move |headers: HeaderMap, upgrade: WebSocketUpgrade| {
let config = config.clone();
let handler = handler.clone();
ws_dispatch(config, handler, headers, upgrade)
})
}
pub(crate) async fn ws_accept<F, Fut>(
config: WsConfig,
headers: HeaderMap,
upgrade: WebSocketUpgrade,
handler: F,
) -> Response
where
F: Fn(WsSession) -> Fut + Clone + Send + Sync + 'static,
Fut: std::future::Future<Output = ()> + Send + 'static,
{
let handler: Arc<dyn Fn(WsSession) -> WsHandlerFuture + Send + Sync + 'static> =
Arc::new(move |session| Box::pin(handler(session)) as _);
ws_dispatch(Arc::new(config), handler, headers, upgrade).await
}
type WsHandlerFuture = std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send>>;
async fn ws_dispatch(
config: Arc<WsConfig>,
handler: Arc<dyn Fn(WsSession) -> WsHandlerFuture + Send + Sync + 'static>,
headers: HeaderMap,
upgrade: WebSocketUpgrade,
) -> Response {
let negotiated = negotiate_subprotocol(&headers, &config.subprotocols);
let mut upgrade = upgrade.max_message_size(config.max_message_bytes);
if let Some(subprotocol) = negotiated.clone() {
upgrade = upgrade.protocols([subprotocol]);
}
upgrade.on_upgrade(move |socket| async move {
let (sink, stream) = socket.split();
let sink = Arc::new(Mutex::new(sink));
let session = WsSession {
sink: sink.clone(),
stream: Mutex::new(stream),
negotiated_subprotocol: negotiated,
max_message_bytes: config.max_message_bytes,
};
let ping_task = config.idle_ping.map(|interval| {
let ping_sink = sink.clone();
tokio::spawn(async move {
let mut ticker = tokio::time::interval(interval);
ticker.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
ticker.tick().await; loop {
ticker.tick().await;
if send_message(&ping_sink, Message::Ping(Default::default()))
.await
.is_err()
{
break;
}
}
})
});
handler(session).await;
if let Some(task) = ping_task {
task.abort();
}
})
}
fn negotiate_subprotocol(headers: &HeaderMap, offered: &[String]) -> Option<String> {
if offered.is_empty() {
return None;
}
let raw = headers.get("sec-websocket-protocol")?.to_str().ok()?;
for client_choice in raw.split(',').map(str::trim) {
if let Some(matched) = offered.iter().find(|name| name.as_str() == client_choice) {
return Some(matched.clone());
}
}
None
}
#[cfg(test)]
mod tests {
use super::*;
use axum::http::HeaderValue;
use axum::Router;
use futures::{SinkExt, StreamExt};
use std::net::SocketAddr;
use tokio::net::TcpListener;
use tokio::sync::Notify;
use tokio_tungstenite::tungstenite::protocol::Message as TungMessage;
async fn echo(session: WsSession) {
while let Ok(Some(message)) = session.recv().await {
match message {
WsMessage::Text(text) => {
if session.send(text).await.is_err() {
break;
}
}
WsMessage::Binary(bytes) => {
if session.send_binary(bytes).await.is_err() {
break;
}
}
}
}
}
async fn spawn_server() -> SocketAddr {
let app = Router::new().route("/ws", ws_route(echo, WsConfig::default()));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
addr
}
#[tokio::test]
async fn ws_echo_roundtrip() {
let addr = spawn_server().await;
let url = format!("ws://{addr}/ws");
let (mut socket, _response) = tokio_tungstenite::connect_async(url).await.unwrap();
socket
.send(TungMessage::Text("hello".into()))
.await
.unwrap();
let echoed = socket.next().await.unwrap().unwrap();
assert_eq!(echoed, TungMessage::Text("hello".into()));
socket
.send(TungMessage::Binary(vec![1, 2, 3, 4].into()))
.await
.unwrap();
let echoed = socket.next().await.unwrap().unwrap();
assert_eq!(echoed, TungMessage::Binary(vec![1, 2, 3, 4].into()));
socket.send(TungMessage::Close(None)).await.unwrap();
}
#[tokio::test]
async fn ws_session_send_does_not_block_on_pending_recv() {
let receive_path_held = std::sync::Arc::new(Notify::new());
let handler_receive_path_held = receive_path_held.clone();
let app = Router::new().route(
"/ws",
ws_route(
move |session| {
let receive_path_held = handler_receive_path_held.clone();
async move {
let _receive_guard = session.stream.lock().await;
receive_path_held.notify_one();
session.send("server-ready").await.unwrap();
}
},
WsConfig::default(),
),
);
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let url = format!("ws://{addr}/ws");
let (mut socket, _response) = tokio_tungstenite::connect_async(url).await.unwrap();
receive_path_held.notified().await;
let message = socket.next().await.unwrap().unwrap();
assert_eq!(message, TungMessage::Text("server-ready".into()));
}
#[tokio::test]
async fn ws_subprotocol_negotiation_picks_first_match() {
let config = WsConfig {
subprotocols: vec!["v1.harn".into(), "v2.harn".into()],
..WsConfig::default()
};
let app = Router::new().route("/ws", ws_route(echo, config));
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
let url = format!("ws://{addr}/ws");
let request =
tokio_tungstenite::tungstenite::client::IntoClientRequest::into_client_request(url)
.unwrap();
let mut request = request;
request.headers_mut().insert(
"Sec-WebSocket-Protocol",
"unsupported, v2.harn".parse().unwrap(),
);
let (mut socket, response) = tokio_tungstenite::connect_async(request).await.unwrap();
assert_eq!(
response
.headers()
.get("sec-websocket-protocol")
.map(|v| v.to_str().unwrap()),
Some("v2.harn"),
);
socket.send(TungMessage::Close(None)).await.unwrap();
}
#[test]
fn negotiate_subprotocol_picks_first_match() {
let mut headers = HeaderMap::new();
headers.insert(
"sec-websocket-protocol",
HeaderValue::from_static("a, b, c"),
);
assert_eq!(
negotiate_subprotocol(&headers, &["c".into(), "a".into()]),
Some("a".into())
);
}
#[test]
fn negotiate_subprotocol_returns_none_when_no_overlap() {
let mut headers = HeaderMap::new();
headers.insert("sec-websocket-protocol", HeaderValue::from_static("x, y"));
assert_eq!(
negotiate_subprotocol(&headers, &["a".into(), "b".into()]),
None
);
}
#[test]
fn negotiate_subprotocol_returns_none_when_offered_empty() {
let mut headers = HeaderMap::new();
headers.insert("sec-websocket-protocol", HeaderValue::from_static("a"));
assert_eq!(negotiate_subprotocol(&headers, &[]), None);
}
}