use std::sync::Arc;
use futures_util::{
SinkExt, StreamExt, TryStreamExt,
stream::{SplitSink, SplitStream},
};
use http::{HeaderValue, Request, Response, StatusCode, request::Parts};
use smallvec::smallvec;
use tokio::io::{AsyncRead, AsyncWrite};
use tokio_tungstenite::{
WebSocketStream,
tungstenite::{
self, Message,
handshake::derive_accept_key,
protocol::{Role, WebSocketConfig},
},
};
use tokio_util::future::FutureExt as _;
use engineioxide_core::{Packet, ProtocolVersion, Sid, Str, TransportType};
use crate::{
DisconnectReason, Socket, body::ResponseBody, config::EngineIoConfig, engine::EngineIo,
errors::Error, handler::EngineIoHandler, socket::InternalRx, transport::make_open_packet,
};
fn ws_response<B>(ws_key: &HeaderValue) -> Response<ResponseBody<B>> {
let derived = derive_accept_key(ws_key.as_bytes());
let sec = derived.parse::<HeaderValue>().unwrap();
Response::builder()
.status(StatusCode::SWITCHING_PROTOCOLS)
.header(http::header::UPGRADE, HeaderValue::from_static("websocket"))
.header(
http::header::CONNECTION,
HeaderValue::from_static("Upgrade"),
)
.header(http::header::SEC_WEBSOCKET_ACCEPT, sec)
.body(ResponseBody::empty_response())
.unwrap() }
pub fn new_req<R: Send + 'static, B, H: EngineIoHandler>(
engine: Arc<EngineIo<H>>,
protocol: ProtocolVersion,
sid: Option<Sid>,
req: Request<R>,
) -> Result<Response<ResponseBody<B>>, Error> {
let (parts, body) = req.into_parts();
let req = Request::from_parts(parts.clone(), body);
let ws_key = parts
.headers
.get("Sec-WebSocket-Key")
.ok_or(Error::InvalidWebSocketKey)?
.clone();
tokio::spawn(async move {
let conn = hyper::upgrade::on(req)
.await
.map(hyper_util::rt::TokioIo::new);
let res = match conn {
Ok(conn) => on_init(engine.clone(), conn, protocol, sid, parts).await,
Err(_e) => {
#[cfg(feature = "tracing")]
tracing::debug!("ws upgrade error: {}", _e);
return;
}
};
match res {
Ok(_) => {
#[cfg(feature = "tracing")]
tracing::debug!(?sid, "ws closed")
}
Err(Error::MultipleWebsocketRequests) => {}
Err(_e) => {
#[cfg(feature = "tracing")]
tracing::debug!(?sid, "ws closed with error: {_e}");
if let Some(sid) = sid {
engine.close_session(sid, DisconnectReason::TransportError);
}
}
}
});
Ok(ws_response(&ws_key))
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", name = "init_websocket", skip(engine, conn, req_data))
)]
pub async fn on_init<H: EngineIoHandler, S>(
engine: Arc<EngineIo<H>>,
conn: S,
protocol: ProtocolVersion,
sid: Option<Sid>,
req_data: Parts,
) -> Result<(), Error>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let ws_config = WebSocketConfig::default().read_buffer_size(engine.config.ws_read_buffer_size);
let ws_init = move || WebSocketStream::from_raw_socket(conn, Role::Server, Some(ws_config));
let (socket, ws) = if let Some(sid) = sid {
match engine.get_socket(sid) {
None => return Err(Error::UnknownSessionID(sid)),
Some(socket) if socket.is_ws() => return Err(Error::MultipleWebsocketRequests),
Some(socket) => {
let mut ws = ws_init().await;
let upgrade_fut = upgrade_handshake::<H, S>(&socket, &mut ws);
match tokio::time::timeout(engine.config.upgrade_timeout, upgrade_fut)
.with_cancellation_token(&socket.cancellation_token)
.await
{
Some(Ok(res)) => res?,
Some(Err(_)) => {
#[cfg(feature = "tracing")]
tracing::debug!(?sid, "ws upgrade timed out, closing session");
engine.close_session(socket.id, DisconnectReason::TransportError);
return Err(Error::Upgrade);
}
None => {
#[cfg(feature = "tracing")]
tracing::debug!(?sid, "socket is being closed");
return Err(Error::Upgrade);
}
}
(socket, ws)
}
}
} else {
let socket = engine.create_session(protocol, TransportType::Websocket, req_data, false);
#[cfg(feature = "tracing")]
tracing::debug!("new websocket connection");
let mut ws = ws_init().await;
init_handshake(socket.id, &mut ws, &engine.config).await?;
socket
.clone()
.spawn_heartbeat(engine.config.ping_interval, engine.config.ping_timeout);
(socket, ws)
};
let (tx, rx) = ws.split();
tokio::spawn(forward_to_socket::<H, S>(socket.clone(), tx));
match forward_to_handler(&engine, rx, &socket)
.with_cancellation_token(&socket.cancellation_token)
.await
{
Some(Err(ref e)) => {
let reason =
Option::<DisconnectReason>::from(e).unwrap_or(DisconnectReason::TransportError);
#[cfg(feature = "tracing")]
tracing::debug!("error when handling packet: {:?}", e);
engine.close_session(socket.id, reason);
}
Some(Ok(())) => {
#[cfg(feature = "tracing")]
tracing::debug!(sid = %socket.id, "ws transport was closed");
engine.close_session(socket.id, DisconnectReason::TransportClose);
}
None => {
#[cfg(feature = "tracing")]
tracing::debug!(sid = %socket.id, "socket is closing");
}
}
Ok(())
}
async fn forward_to_handler<H: EngineIoHandler, S>(
engine: &Arc<EngineIo<H>>,
mut rx: SplitStream<WebSocketStream<S>>,
socket: &Arc<Socket<H::Data>>,
) -> Result<(), Error>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
while let Some(msg) = rx.try_next().await? {
match msg {
Message::Text(msg) => match Packet::parse(socket.protocol, ws_bytes_to_str(msg))? {
Packet::Close => {
#[cfg(feature = "tracing")]
tracing::debug!("[sid={}] closing session", socket.id);
engine.close_session(socket.id, DisconnectReason::TransportClose);
break;
}
Packet::Pong | Packet::Ping => socket
.heartbeat_tx
.try_send(())
.map_err(|_| Error::HeartbeatTimeout),
Packet::Message(msg) => {
engine.handler.on_message(msg, socket.clone());
Ok(())
}
p => return Err(Error::BadPacket(p)),
},
Message::Binary(mut data) => {
if socket.protocol == ProtocolVersion::V3 && !data.is_empty() {
data = data.split_off(1);
}
engine.handler.on_binary(data, socket.clone());
Ok(())
}
Message::Close(_) => {
#[cfg(feature = "tracing")]
tracing::debug!("websocket closed, closing session");
engine.close_session(socket.id, DisconnectReason::TransportClose);
break;
}
_ => {
#[cfg(feature = "tracing")]
tracing::debug!(sid = ?socket.id, "unexpected ws message");
Ok(())
}
}?
}
Ok(())
}
async fn forward_to_socket<H: EngineIoHandler, S>(
socket: Arc<Socket<H::Data>>,
mut tx: SplitSink<WebSocketStream<S>, Message>,
) where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let mut internal_rx = socket.internal_rx.try_lock().unwrap();
let InternalRx {
buffered_rx,
volatile_rx,
..
} = &mut *internal_rx;
macro_rules! map_fn {
($item:ident) => {
let res = match $item {
Packet::Binary(bin) | Packet::BinaryV3(bin) => {
if socket.protocol == ProtocolVersion::V3 {
let mut buff = Vec::with_capacity(bin.len() + 1);
buff.push(0x04);
buff.extend(bin);
tx.feed(Message::Binary(buff.into())).await
} else {
tx.feed(Message::Binary(bin)).await
}
}
Packet::Noop => Ok(()),
_ => {
let packet: String = $item.try_into().unwrap();
tx.feed(Message::Text(packet.into())).await
}
};
if let Err(_e) = res {
#[cfg(feature = "tracing")]
tracing::debug!("[sid={}] error sending packet: {}", socket.id, _e);
}
};
}
loop {
tokio::select! {
biased;
items = buffered_rx.recv() => {
match items {
Some(packets) => {
for item in packets {
map_fn!(item);
}
while let Ok(packets) = buffered_rx.try_recv() {
for item in packets {
map_fn!(item);
}
}
}
None => break
}
}
_ = socket.cancellation_token.cancelled() => break,
Ok(()) = volatile_rx.changed() => {
let val = volatile_rx.borrow_and_update().clone();
if let Some(packets) = val {
#[cfg(feature = "tracing")]
tracing::debug!(sid = ?socket.id, "ws volatile flush: {} packets", packets.len());
for item in packets {
map_fn!(item);
}
}
},
}
#[cfg(feature = "tracing")]
tracing::trace!(sid = %socket.id, "ws flush");
tx.flush().await.ok();
}
#[cfg(feature = "tracing")]
tracing::trace!(sid = %socket.id, "ws closing flush");
tx.flush().await.ok();
tx.close().await.ok();
}
async fn init_handshake<S>(
sid: Sid,
ws: &mut WebSocketStream<S>,
config: &EngineIoConfig,
) -> Result<(), Error>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
let packet = Packet::Open(make_open_packet(TransportType::Websocket, sid, config));
let packet: String = packet.into();
ws.send(Message::Text(packet.into())).await?;
Ok(())
}
#[cfg_attr(
feature = "tracing",
tracing::instrument(level = "trace", skip(socket, ws))
)]
async fn upgrade_handshake<H: EngineIoHandler, S>(
socket: &Arc<Socket<H::Data>>,
ws: &mut WebSocketStream<S>,
) -> Result<(), Error>
where
S: AsyncRead + AsyncWrite + Unpin + Send + 'static,
{
#[cfg(feature = "tracing")]
tracing::debug!("starting websocket connection upgrade");
socket.start_upgrade();
socket
.internal_tx
.send(smallvec![Packet::Noop])
.await
.map_err(|_| Error::Upgrade)?;
let _lock = socket.internal_rx.lock().await;
let msg = match ws.next().await {
Some(Ok(Message::Text(d))) => d,
_ => Err(Error::Upgrade)?,
};
match Packet::parse(socket.protocol, ws_bytes_to_str(msg))? {
Packet::PingUpgrade => {
#[cfg(feature = "tracing")]
tracing::debug!("received first ping upgrade");
ws.send(Message::Text(String::from(Packet::PongUpgrade).into()))
.await?;
}
p => Err(Error::BadPacket(p))?,
};
let msg = match ws.next().await {
Some(Ok(Message::Text(d))) => d,
Some(Ok(Message::Close(_))) => {
#[cfg(feature = "tracing")]
tracing::debug!("ws stream closed before upgrade");
Err(Error::Upgrade)?
}
_ => {
#[cfg(feature = "tracing")]
tracing::debug!("unexpected ws message before upgrade");
Err(Error::Upgrade)?
}
};
match Packet::parse(socket.protocol, ws_bytes_to_str(msg))? {
Packet::Upgrade => {
#[cfg(feature = "tracing")]
tracing::debug!("ws upgraded successfully")
}
p => Err(Error::BadPacket(p))?,
};
socket.upgrade_to_websocket();
Ok(())
}
fn ws_bytes_to_str(bytes: tungstenite::Utf8Bytes) -> Str {
unsafe { Str::from_bytes_unchecked(bytes.into()) }
}