use std::sync::Arc;
use futures_util::{
SinkExt, StreamExt, TryStreamExt,
stream::{SplitSink, SplitStream},
};
use http::{HeaderValue, Request, Response, StatusCode, request::Parts};
use tokio::io::{AsyncRead, AsyncWrite};
use tokio_tungstenite::{
WebSocketStream,
tungstenite::{
Message,
handshake::derive_accept_key,
protocol::{Role, WebSocketConfig},
},
};
use engineioxide_core::Sid;
use crate::{
DisconnectReason, Socket,
body::ResponseBody,
config::EngineIoConfig,
engine::EngineIo,
errors::Error,
handler::EngineIoHandler,
packet::{OpenPacket, Packet},
service::ProtocolVersion,
service::TransportType,
};
fn ws_response<B>(ws_key: &HeaderValue) -> Result<Response<ResponseBody<B>>, http::Error> {
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())
}
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::HttpErrorResponse(StatusCode::BAD_REQUEST))?
.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;
upgrade_handshake::<H, S>(&socket, &mut ws).await?;
(socket, ws)
}
}
} else {
let socket = engine.create_session(
protocol,
TransportType::Websocket,
req_data,
#[cfg(feature = "v3")]
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();
let cancel_token = socket.cancellation_token.clone();
tokio::spawn(
cancel_token
.clone()
.run_until_cancelled_owned(forward_to_socket::<H, S>(socket.clone(), tx)),
);
cancel_token
.run_until_cancelled_owned(async move {
if let Err(ref e) = forward_to_handler(&engine, rx, &socket).await {
#[cfg(feature = "tracing")]
tracing::debug!("error when handling packet: {:?}", e);
if let Some(reason) = e.into() {
engine.close_session(socket.id, reason);
}
} else {
engine.close_session(socket.id, DisconnectReason::TransportClose);
}
})
.await;
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::try_from(msg)? {
Packet::Close => 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();
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::Close => {
tx.send(Message::Close(None)).await.ok();
internal_rx.close();
break;
},
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);
}
};
}
while let Some(items) = internal_rx.recv().await {
for item in items {
map_fn!(item);
}
while let Ok(items) = internal_rx.try_recv() {
for item in items {
map_fn!(item);
}
}
tx.flush().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(OpenPacket::new(TransportType::Websocket, sid, config));
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.send(Packet::Noop)?;
let _lock = socket.internal_rx.lock().await;
let msg = match ws.next().await {
Some(Ok(Message::Text(d))) => d,
_ => Err(Error::Upgrade)?,
};
match Packet::try_from(msg)? {
Packet::PingUpgrade => {
#[cfg(feature = "tracing")]
tracing::debug!("received first ping upgrade");
ws.send(Message::Text(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::try_from(msg)? {
Packet::Upgrade => {
#[cfg(feature = "tracing")]
tracing::debug!("ws upgraded successfully")
}
p => Err(Error::BadPacket(p))?,
};
socket.upgrade_to_websocket();
Ok(())
}