#![allow(clippy::future_not_send)]
use crate::{
header,
websocket::{
ffi,
session::{internal_error_frame, IntoWebSocketOutcome},
types::{
offered_protocols, select_protocol, WebSocketCloseFrame, WebSocketError,
WebSocketResult,
},
},
Method, Request, Response, StatusCode,
};
use futures_channel::mpsc::{self, UnboundedReceiver, UnboundedSender};
use futures_core::Stream;
use http_kit::utils::ByteStr;
use http_kit::ws::{WebSocketConfig, WebSocketMessage};
use serde::Serialize;
use skyzen_core::{error::ErrorChain, Extractor, Responder};
use std::{
cell::RefCell,
future::{ready, Future},
pin::Pin,
rc::Rc,
task::{Context, Poll},
};
use wasm_bindgen::{prelude::*, JsCast};
const fn ensure_within_limit(config: &WebSocketConfig, len: usize) -> WebSocketResult<()> {
if let Some(limit) = config.max_message_size {
if len > limit {
return Err(WebSocketError::MessageTooLarge { len, limit });
}
}
Ok(())
}
fn close_socket(
socket: &ffi::WebSocket,
close_frame: Option<WebSocketCloseFrame>,
) -> WebSocketResult<()> {
let (code, reason) = close_frame.map_or((None, None), |frame| (Some(frame.code), Some(frame)));
socket
.close(code, reason.as_ref().map(|frame| frame.reason.as_str()))
.map_err(|error| WebSocketError::Protocol(format!("{error:?}")))
}
pub struct WebSocket {
inner: ffi::WebSocket,
rx: UnboundedReceiver<WebSocketResult<WebSocketMessage>>,
closures: Rc<RefCell<EventClosures>>,
config: WebSocketConfig,
}
#[allow(clippy::struct_field_names)]
struct EventClosures {
_on_message: Closure<dyn FnMut(ffi::MessageEvent)>,
_on_close: Closure<dyn FnMut(ffi::CloseEvent)>,
_on_error: Closure<dyn FnMut(ffi::ErrorEvent)>,
}
impl WebSocket {
pub(crate) fn from_ffi_socket(socket: ffi::WebSocket, config: WebSocketConfig) -> Self {
let (tx, rx) = mpsc::unbounded();
let closures = Self::setup_event_handlers(&socket, tx, &config);
Self {
inner: socket,
rx,
closures: Rc::new(RefCell::new(closures)),
config,
}
}
fn setup_event_handlers(
socket: &ffi::WebSocket,
tx: UnboundedSender<WebSocketResult<WebSocketMessage>>,
config: &WebSocketConfig,
) -> EventClosures {
let tx_message = tx.clone();
let max_message_size = config.max_message_size;
let on_message = Closure::wrap(Box::new(move |event: ffi::MessageEvent| {
let data = event.data();
let message = if let Some(text) = data.as_string() {
WebSocketMessage::Text(text.into())
} else if data.is_instance_of::<js_sys::ArrayBuffer>()
|| data.is_instance_of::<js_sys::Uint8Array>()
{
WebSocketMessage::Binary(js_sys::Uint8Array::new(&data).to_vec().into())
} else {
return;
};
let len = match &message {
WebSocketMessage::Text(text) => text.len(),
WebSocketMessage::Binary(bytes) => bytes.len(),
_ => 0,
};
if let Some(limit) = max_message_size {
if len > limit {
let _ = tx_message
.unbounded_send(Err(WebSocketError::MessageTooLarge { len, limit }));
return;
}
}
let _ = tx_message.unbounded_send(Ok(message));
}) as Box<dyn FnMut(ffi::MessageEvent)>);
let tx_close = tx.clone();
let on_close = Closure::wrap(Box::new(move |event: ffi::CloseEvent| {
tracing::debug!(
code = event.code(),
reason = %event.reason(),
was_clean = event.was_clean(),
"websocket closed by peer"
);
let _ = tx_close.unbounded_send(Ok(WebSocketMessage::Close));
tx_close.close_channel();
}) as Box<dyn FnMut(ffi::CloseEvent)>);
let on_error = Closure::wrap(Box::new(move |event: ffi::ErrorEvent| {
let _ = tx.unbounded_send(Err(WebSocketError::Protocol(event.message())));
tx.close_channel();
}) as Box<dyn FnMut(ffi::ErrorEvent)>);
socket.add_event_listener("message", on_message.as_ref().unchecked_ref());
socket.add_event_listener("close", on_close.as_ref().unchecked_ref());
socket.add_event_listener("error", on_error.as_ref().unchecked_ref());
EventClosures {
_on_message: on_message,
_on_close: on_close,
_on_error: on_error,
}
}
#[cfg(feature = "json")]
pub async fn send<T: Serialize>(&mut self, value: T) -> WebSocketResult<()> {
let payload = serde_json::to_string(&value)?;
self.send_text(payload).await
}
pub fn send_text(
&mut self,
text: impl Into<ByteStr>,
) -> impl Future<Output = WebSocketResult<()>> {
let text = text.into();
ready(
ensure_within_limit(&self.config, text.len()).and_then(|()| {
self.inner
.send(&JsValue::from_str(&text))
.map_err(|e| WebSocketError::Protocol(format!("{e:?}")))
}),
)
}
pub fn send_binary(
&mut self,
data: impl Into<Vec<u8>>,
) -> impl Future<Output = WebSocketResult<()>> {
let bytes = data.into();
ready(
ensure_within_limit(&self.config, bytes.len()).and_then(|()| {
let array = js_sys::Uint8Array::from(&bytes[..]);
self.inner
.send(&array.into())
.map_err(|e| WebSocketError::Protocol(format!("{e:?}")))
}),
)
}
pub fn send_ping(
&mut self,
_data: impl Into<Vec<u8>>,
) -> impl Future<Output = WebSocketResult<()>> {
ready(Err(WebSocketError::Protocol(
"Ping frames not supported on WASM platform".into(),
)))
}
pub fn send_pong(
&mut self,
_data: impl Into<Vec<u8>>,
) -> impl Future<Output = WebSocketResult<()>> {
ready(Err(WebSocketError::Protocol(
"Pong frames not supported on WASM platform".into(),
)))
}
pub async fn send_message(&mut self, message: WebSocketMessage) -> WebSocketResult<()> {
match message {
WebSocketMessage::Text(text) => self.send_text(text).await,
WebSocketMessage::Binary(data) => self.send_binary(data).await,
WebSocketMessage::Close => self.close(None).await,
WebSocketMessage::Ping(_) => self.send_ping(vec![]).await,
WebSocketMessage::Pong(_) => self.send_pong(vec![]).await,
}
}
#[cfg(feature = "json")]
pub async fn recv_json<T: serde::de::DeserializeOwned>(
&mut self,
) -> Option<WebSocketResult<T>> {
use futures_util::StreamExt;
loop {
match self.next().await {
Some(Ok(msg)) => {
if let Some(result) = msg.into_json() {
return Some(result.map_err(WebSocketError::from));
}
}
Some(Err(e)) => return Some(Err(e)),
None => return None,
}
}
}
#[must_use]
pub const fn get_config(&self) -> &WebSocketConfig {
&self.config
}
pub fn close(
&mut self,
close_frame: Option<WebSocketCloseFrame>,
) -> impl Future<Output = WebSocketResult<()>> {
ready(close_socket(&self.inner, close_frame))
}
#[must_use]
pub fn split(self) -> (WebSocketSender, WebSocketReceiver) {
(
WebSocketSender {
inner: self.inner,
config: self.config.clone(),
_closures: self.closures.clone(),
},
WebSocketReceiver {
rx: self.rx,
config: self.config,
_closures: self.closures,
},
)
}
}
impl std::fmt::Debug for WebSocket {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocket").finish_non_exhaustive()
}
}
impl Stream for WebSocket {
type Item = WebSocketResult<WebSocketMessage>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Pin::new(&mut self.rx).poll_next(cx)
}
}
pub struct WebSocketSender {
inner: ffi::WebSocket,
config: WebSocketConfig,
_closures: Rc<RefCell<EventClosures>>,
}
impl WebSocketSender {
#[cfg(feature = "json")]
pub async fn send<T: Serialize>(&mut self, value: T) -> WebSocketResult<()> {
let payload = serde_json::to_string(&value)?;
self.send_text(payload).await
}
pub fn send_text(
&mut self,
text: impl Into<ByteStr>,
) -> impl Future<Output = WebSocketResult<()>> {
let text = text.into();
ready(
ensure_within_limit(&self.config, text.len()).and_then(|()| {
self.inner
.send(&JsValue::from_str(&text))
.map_err(|e| WebSocketError::Protocol(format!("{e:?}")))
}),
)
}
pub fn send_binary(
&mut self,
data: impl Into<Vec<u8>>,
) -> impl Future<Output = WebSocketResult<()>> {
let bytes = data.into();
ready(
ensure_within_limit(&self.config, bytes.len()).and_then(|()| {
let array = js_sys::Uint8Array::from(&bytes[..]);
self.inner
.send(&array.into())
.map_err(|e| WebSocketError::Protocol(format!("{e:?}")))
}),
)
}
pub fn send_ping(
&mut self,
_data: impl Into<Vec<u8>>,
) -> impl Future<Output = WebSocketResult<()>> {
ready(Err(WebSocketError::Protocol(
"Ping frames not supported on WASM platform".into(),
)))
}
pub fn send_pong(
&mut self,
_data: impl Into<Vec<u8>>,
) -> impl Future<Output = WebSocketResult<()>> {
ready(Err(WebSocketError::Protocol(
"Pong frames not supported on WASM platform".into(),
)))
}
pub async fn send_message(&mut self, message: WebSocketMessage) -> WebSocketResult<()> {
match message {
WebSocketMessage::Text(text) => self.send_text(text).await,
WebSocketMessage::Binary(data) => self.send_binary(data).await,
WebSocketMessage::Close => self.close(None).await,
WebSocketMessage::Ping(_) | WebSocketMessage::Pong(_) => Err(WebSocketError::Protocol(
"Ping/Pong not supported on WASM".into(),
)),
}
}
pub fn close(
&mut self,
close_frame: Option<WebSocketCloseFrame>,
) -> impl Future<Output = WebSocketResult<()>> {
ready(close_socket(&self.inner, close_frame))
}
#[must_use]
pub const fn get_config(&self) -> &WebSocketConfig {
&self.config
}
}
impl std::fmt::Debug for WebSocketSender {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocketSender").finish_non_exhaustive()
}
}
pub struct WebSocketReceiver {
rx: UnboundedReceiver<WebSocketResult<WebSocketMessage>>,
config: WebSocketConfig,
_closures: Rc<RefCell<EventClosures>>,
}
impl WebSocketReceiver {
#[cfg(feature = "json")]
pub async fn recv_json<T: serde::de::DeserializeOwned>(
&mut self,
) -> Option<WebSocketResult<T>> {
use futures_util::StreamExt;
loop {
match self.next().await {
Some(Ok(msg)) => {
if let Some(result) = msg.into_json() {
return Some(result.map_err(WebSocketError::from));
}
}
Some(Err(e)) => return Some(Err(e)),
None => return None,
}
}
}
#[must_use]
pub const fn get_config(&self) -> &WebSocketConfig {
&self.config
}
}
impl std::fmt::Debug for WebSocketReceiver {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocketReceiver").finish_non_exhaustive()
}
}
impl Stream for WebSocketReceiver {
type Item = WebSocketResult<WebSocketMessage>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
Pin::new(&mut self.rx).poll_next(cx)
}
}
#[skyzen::error(status = StatusCode::BAD_REQUEST)]
pub enum WebSocketUpgradeError {
#[error("Method not allowed", status = StatusCode::METHOD_NOT_ALLOWED)]
MethodNotAllowed,
#[error("Missing or invalid upgrade header")]
MissingUpgradeHeader,
#[error("Missing Connection header for WebSocket request")]
MissingConnectionHeader,
#[error("Missing Sec-WebSocket-Key header")]
MissingSecWebSocketKey,
#[error("Upgrade header must be `websocket`")]
InvalidUpgradeHeader,
#[error("Invalid Connection header for WebSocket request")]
InvalidConnectionHeader,
#[error("Unsupported Sec-WebSocket-Version. Only version 13 is accepted")]
UnsupportedVersion,
}
struct SendSyncWebSocketPair(ffi::WebSocketPair);
unsafe impl Send for SendSyncWebSocketPair {}
unsafe impl Sync for SendSyncWebSocketPair {}
pub struct WebSocketUpgrade {
pair: SendSyncWebSocketPair,
requested_protocols: Vec<String>,
response_protocol: Option<header::HeaderValue>,
config: WebSocketConfig,
}
impl WebSocketUpgrade {
fn new(requested_protocols: Vec<String>) -> Self {
Self {
pair: SendSyncWebSocketPair(ffi::WebSocketPair::new()),
requested_protocols,
response_protocol: None,
config: WebSocketConfig::default(),
}
}
#[must_use]
pub fn protocols<I, S>(mut self, protocols: I) -> Self
where
I: IntoIterator<Item = S>,
S: AsRef<str>,
{
let supported: Vec<String> = protocols
.into_iter()
.map(|protocol| protocol.as_ref().to_owned())
.collect();
self.response_protocol = select_protocol(&self.requested_protocols, &supported);
self
}
#[must_use]
pub fn protocol(mut self, protocol: header::HeaderValue) -> Self {
self.response_protocol = Some(protocol);
self
}
#[must_use]
pub fn requested_protocols(&self) -> &[String] {
&self.requested_protocols
}
#[must_use]
pub const fn config(mut self, config: WebSocketConfig) -> Self {
self.config = config;
self
}
#[must_use]
pub const fn max_message_size(mut self, max_size: Option<usize>) -> Self {
self.config.max_message_size = max_size;
self
}
pub fn on_upgrade<F, Fut, R>(self, callback: F) -> WebSocketUpgradeResponder
where
F: FnOnce(WebSocket) -> Fut + 'static,
Fut: std::future::Future<Output = R> + 'static,
R: IntoWebSocketOutcome + 'static,
{
let Self {
pair,
response_protocol,
config,
requested_protocols: _,
} = self;
let pair = pair.0;
let server = pair.server();
let client = pair.client();
server.accept();
let closer = server.clone();
let socket = WebSocket::from_ffi_socket(server, config);
wasm_bindgen_futures::spawn_local(async move {
if let Err(error) = callback(socket).await.into_outcome() {
tracing::error!(error = %ErrorChain(&error), "websocket session handler failed");
if let Err(error) = close_socket(&closer, Some(internal_error_frame())) {
tracing::debug!("failed to close a failed websocket session: {error}");
}
}
});
WebSocketUpgradeResponder {
client: SendSyncWebSocket(client),
response_protocol,
}
}
}
impl std::fmt::Debug for WebSocketUpgrade {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocketUpgrade")
.field("requested_protocols", &self.requested_protocols)
.field("response_protocol", &self.response_protocol)
.field("config", &self.config)
.finish_non_exhaustive()
}
}
fn header_has_token(value: &header::HeaderValue, token: &str) -> bool {
value.to_str().is_ok_and(|value| {
value
.split(',')
.any(|part| part.trim().eq_ignore_ascii_case(token))
})
}
impl Extractor for WebSocketUpgrade {
type Error = WebSocketUpgradeError;
fn extract(request: &mut Request) -> impl Future<Output = Result<Self, Self::Error>> + Send {
ready(validate_upgrade(request))
}
}
fn validate_upgrade(request: &Request) -> Result<WebSocketUpgrade, WebSocketUpgradeError> {
if request.method() != Method::GET {
return Err(WebSocketUpgradeError::MethodNotAllowed);
}
let headers = request.headers();
headers
.get(header::SEC_WEBSOCKET_KEY)
.ok_or(WebSocketUpgradeError::MissingSecWebSocketKey)?;
let connection = headers
.get(header::CONNECTION)
.ok_or(WebSocketUpgradeError::MissingConnectionHeader)?;
if !header_has_token(connection, "upgrade") {
return Err(WebSocketUpgradeError::InvalidConnectionHeader);
}
let upgrade_header = headers
.get(header::UPGRADE)
.ok_or(WebSocketUpgradeError::MissingUpgradeHeader)?;
if !upgrade_header
.to_str()
.is_ok_and(|value| value.eq_ignore_ascii_case("websocket"))
{
return Err(WebSocketUpgradeError::InvalidUpgradeHeader);
}
match headers.get(header::SEC_WEBSOCKET_VERSION) {
Some(version) if version == "13" => {}
_ => return Err(WebSocketUpgradeError::UnsupportedVersion),
}
Ok(WebSocketUpgrade::new(offered_protocols(headers)))
}
#[derive(Clone)]
pub struct SendSyncWebSocket(pub(crate) ffi::WebSocket);
impl SendSyncWebSocket {
#[must_use]
pub fn into_inner(self) -> ffi::WebSocket {
self.0
}
}
impl std::fmt::Debug for SendSyncWebSocket {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SendSyncWebSocket").finish_non_exhaustive()
}
}
unsafe impl Send for SendSyncWebSocket {}
unsafe impl Sync for SendSyncWebSocket {}
pub struct WebSocketUpgradeResponder {
client: SendSyncWebSocket,
response_protocol: Option<header::HeaderValue>,
}
impl std::fmt::Debug for WebSocketUpgradeResponder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocketUpgradeResponder")
.field("response_protocol", &self.response_protocol)
.finish_non_exhaustive()
}
}
impl Responder for WebSocketUpgradeResponder {
type Error = std::convert::Infallible;
fn respond_to(self, _request: &Request, response: &mut Response) -> Result<(), Self::Error> {
*response.status_mut() = StatusCode::SWITCHING_PROTOCOLS;
if let Some(protocol) = self.response_protocol {
response
.headers_mut()
.insert(header::SEC_WEBSOCKET_PROTOCOL, protocol);
}
response.extensions_mut().insert(self.client);
Ok(())
}
}