use super::compat::{AllowStd, Direction};
use super::deflate::PerMessageDeflate;
use crate::http::error::Error;
use bytes::{Bytes, BytesMut};
use futures_util::{Sink, SinkExt, Stream};
use hyper::header::HeaderValue;
use hyper_util::rt::TokioIo;
use std::collections::VecDeque;
use std::future::poll_fn;
use std::pin::Pin;
use std::task::{Context, Poll};
use tokio_util::sync::PollSender;
use tungstenite::Error as WsError;
use tungstenite::protocol::WebSocketConfig;
use tungstenite::protocol::frame::coding::{Control, Data as OpData, OpCode};
use tungstenite::protocol::frame::{CloseFrame, Frame, FrameSocket, Utf8Bytes};
pub use tungstenite::Message;
type Io = AllowStd<TokioIo<hyper::upgrade::Upgraded>>;
const SPLIT_CHANNEL_CAPACITY: usize = 32;
pub struct WebSocket {
frames: FrameSocket<Io>,
protocol: Option<HeaderValue>,
config: WebSocketConfig,
deflate: Option<PerMessageDeflate>,
fragment: Option<Fragment>,
outgoing: VecDeque<Frame>,
flush_needed: bool,
sent_close: bool,
closed: bool,
}
struct Fragment {
opcode: OpData,
compressed: bool,
buffer: BytesMut,
}
impl std::fmt::Debug for WebSocket {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("WebSocket").finish_non_exhaustive()
}
}
fn protocol_error(err: &WsError) -> Error {
Error::Internal(err.to_string())
}
fn unmask(data: &mut [u8], mask: [u8; 4]) {
let mask8 = u64::from_ne_bytes([
mask[0], mask[1], mask[2], mask[3], mask[0], mask[1], mask[2], mask[3],
]);
let mut chunks8 = data.chunks_exact_mut(8);
for chunk in &mut chunks8 {
let word =
u64::from_ne_bytes([
chunk[0], chunk[1], chunk[2], chunk[3], chunk[4], chunk[5], chunk[6], chunk[7],
]) ^ mask8;
chunk.copy_from_slice(&word.to_ne_bytes());
}
let mask4 = u32::from_ne_bytes(mask);
let mut chunks4 = chunks8.into_remainder().chunks_exact_mut(4);
for chunk in &mut chunks4 {
let word = u32::from_ne_bytes([chunk[0], chunk[1], chunk[2], chunk[3]]) ^ mask4;
chunk.copy_from_slice(&word.to_ne_bytes());
}
for (i, byte) in chunks4.into_remainder().iter_mut().enumerate() {
*byte ^= mask[i % 4];
}
}
fn parse_close(payload: &[u8]) -> Result<Option<CloseFrame>, Error> {
match payload.len() {
0 => Ok(None),
1 => Err(Error::Internal("invalid WebSocket close frame".to_string())),
_ => {
let code = u16::from_be_bytes([payload[0], payload[1]]);
let reason = std::str::from_utf8(&payload[2..])
.map_err(|_| Error::Internal("WebSocket close reason is not UTF-8".to_string()))?;
Ok(Some(CloseFrame { code: code.into(), reason: reason.to_string().into() }))
}
}
}
fn poll_io<T>(
frames: &mut FrameSocket<Io>,
direction: &Direction,
cx: &Context<'_>,
f: impl FnOnce(&mut FrameSocket<Io>) -> Result<T, WsError>,
) -> Poll<Result<T, WsError>> {
frames.get_mut().register(direction, cx);
match f(frames) {
Ok(value) => Poll::Ready(Ok(value)),
Err(WsError::Io(err)) if err.kind() == std::io::ErrorKind::WouldBlock => Poll::Pending,
Err(err) => Poll::Ready(Err(err)),
}
}
impl WebSocket {
pub(super) fn new(
io: TokioIo<hyper::upgrade::Upgraded>,
protocol: Option<HeaderValue>,
config: WebSocketConfig,
deflate: Option<PerMessageDeflate>,
) -> Self {
Self {
frames: FrameSocket::new(AllowStd::new(io)),
protocol,
config,
deflate,
fragment: None,
outgoing: VecDeque::new(),
flush_needed: false,
sent_close: false,
closed: false,
}
}
fn poll_drain_outgoing(&mut self, cx: &Context<'_>) -> Poll<Result<(), WsError>> {
while let Some(frame) = self.outgoing.pop_front() {
match poll_io(&mut self.frames, &Direction::Write, cx, |fs| fs.write(frame)) {
Poll::Ready(Ok(())) => self.flush_needed = true,
other => return other,
}
}
if !self.flush_needed {
return Poll::Ready(Ok(()));
}
match poll_io(&mut self.frames, &Direction::Write, cx, FrameSocket::flush) {
Poll::Ready(Ok(())) => {
self.flush_needed = false;
Poll::Ready(Ok(()))
}
other => other,
}
}
fn finish_message(
&mut self,
opcode: OpData,
compressed: bool,
payload: Bytes,
) -> Result<Option<Message>, Error> {
let bytes = if compressed {
let deflate = self.deflate.as_mut().ok_or_else(|| {
Error::Internal("RSV1 set but permessage-deflate was not negotiated".to_string())
})?;
Bytes::from(deflate.decompress(&payload, self.config.max_message_size)?)
} else {
payload
};
match opcode {
OpData::Text => {
let text = Utf8Bytes::try_from(bytes)
.map_err(|e| Error::Internal(format!("invalid UTF-8 in text message: {e}")))?;
Ok(Some(Message::Text(text)))
}
OpData::Binary => Ok(Some(Message::Binary(bytes))),
OpData::Continue | OpData::Reserved(_) => unreachable!("caller only passes Text/Binary"),
}
}
fn handle_frame(&mut self, frame: Frame) -> Result<Option<Message>, Error> {
let header = frame.header().clone();
let raw = frame.into_payload();
let payload = if let Some(mask) = header.mask {
let mut buf = raw.try_into_mut().unwrap_or_else(|shared| BytesMut::from(&shared[..]));
unmask(&mut buf, mask);
buf.freeze()
} else if self.config.accept_unmasked_frames {
raw
} else {
return Err(Error::Internal(
"received an unmasked frame from the client".to_string(),
));
};
match header.opcode {
OpCode::Control(Control::Ping) => {
self.outgoing.push_back(Frame::pong(payload.clone()));
Ok(Some(Message::Ping(payload)))
}
OpCode::Control(Control::Pong) => Ok(Some(Message::Pong(payload))),
OpCode::Control(Control::Close) => {
let close_frame = parse_close(&payload)?;
if !self.sent_close {
self.outgoing.push_back(Frame::close(close_frame.clone()));
self.sent_close = true;
}
self.closed = true;
Ok(Some(Message::Close(close_frame)))
}
OpCode::Control(Control::Reserved(code)) => Err(Error::Internal(format!(
"received reserved WebSocket control opcode {code}"
))),
OpCode::Data(data @ (OpData::Text | OpData::Binary)) => {
if self.fragment.is_some() {
return Err(Error::Internal(
"received a new data frame while a fragmented message was in progress"
.to_string(),
));
}
if header.is_final {
self.finish_message(data, header.rsv1, payload)
} else {
let buffer =
payload.try_into_mut().unwrap_or_else(|shared| BytesMut::from(&shared[..]));
self.fragment = Some(Fragment { opcode: data, compressed: header.rsv1, buffer });
Ok(None)
}
}
OpCode::Data(OpData::Continue) => {
let fragment = self.fragment.as_mut().ok_or_else(|| {
Error::Internal("received a continuation frame with no message in progress".to_string())
})?;
fragment.buffer.extend_from_slice(&payload);
if let Some(max) = self.config.max_message_size
&& fragment.buffer.len() > max
{
return Err(Error::Internal(
"WebSocket message exceeds the configured maximum size".to_string(),
));
}
if header.is_final {
let Fragment { opcode, compressed, buffer } =
self.fragment.take().unwrap_or_else(|| unreachable!());
self.finish_message(opcode, compressed, buffer.freeze())
} else {
Ok(None)
}
}
OpCode::Data(OpData::Reserved(code)) => {
Err(Error::Internal(format!("received reserved WebSocket data opcode {code}")))
}
}
}
fn poll_recv(&mut self, cx: &Context<'_>) -> Poll<Option<Result<Message, Error>>> {
loop {
match self.poll_drain_outgoing(cx) {
Poll::Ready(Ok(())) => {}
Poll::Ready(Err(err)) => {
self.closed = true;
return Poll::Ready(Some(Err(protocol_error(&err))));
}
Poll::Pending => return Poll::Pending,
}
if self.closed {
return Poll::Ready(None);
}
let max_frame_size = self.config.max_frame_size;
let frame = match poll_io(&mut self.frames, &Direction::Read, cx, |fs| {
fs.read(max_frame_size)
}) {
Poll::Ready(Ok(Some(frame))) => frame,
Poll::Ready(Ok(None)) => {
self.closed = true;
return Poll::Ready(None);
}
Poll::Ready(Err(err)) => {
self.closed = true;
return Poll::Ready(Some(Err(protocol_error(&err))));
}
Poll::Pending => return Poll::Pending,
};
match self.handle_frame(frame) {
Ok(Some(msg)) => return Poll::Ready(Some(Ok(msg))),
Ok(None) => {}
Err(err) => {
self.closed = true;
return Poll::Ready(Some(Err(err)));
}
}
}
}
fn queue_data(&mut self, opcode: OpData, payload: Bytes) -> Result<(), Error> {
let compressed = match &mut self.deflate {
Some(deflate) => deflate.compress_if_smaller(&payload)?,
None => None,
};
let (rsv1, bytes) =
compressed.map_or_else(move || (false, payload), |c| (true, Bytes::from(c)));
let mut frame = Frame::message(bytes, OpCode::Data(opcode), true);
frame.header_mut().rsv1 = rsv1;
self.outgoing.push_back(frame);
Ok(())
}
fn queue_message(&mut self, msg: Message) -> Result<(), Error> {
match msg {
Message::Text(text) => self.queue_data(OpData::Text, Bytes::from(text)),
Message::Binary(data) => self.queue_data(OpData::Binary, data),
Message::Ping(data) => {
self.outgoing.push_back(Frame::ping(data));
Ok(())
}
Message::Pong(data) => {
self.outgoing.push_back(Frame::pong(data));
Ok(())
}
Message::Close(frame) => {
self.outgoing.push_back(Frame::close(frame));
self.sent_close = true;
Ok(())
}
Message::Frame(frame) => {
self.outgoing.push_back(frame);
Ok(())
}
}
}
pub async fn recv(&mut self) -> Option<Result<Message, Error>> {
poll_fn(|cx| self.poll_recv(cx)).await
}
pub async fn send(&mut self, msg: Message) -> Result<(), Error> {
self.queue_message(msg)?;
poll_fn(|cx| self.poll_drain_outgoing(cx)).await.map_err(|e| protocol_error(&e))
}
pub async fn flush(&mut self) -> Result<(), Error> {
poll_fn(|cx| self.poll_drain_outgoing(cx)).await.map_err(|e| protocol_error(&e))
}
pub async fn close(mut self) -> Result<(), Error> {
if !self.sent_close {
self.outgoing.push_back(Frame::close(None));
self.sent_close = true;
}
poll_fn(|cx| self.poll_drain_outgoing(cx)).await.map_err(|e| protocol_error(&e))
}
#[must_use]
pub const fn protocol(&self) -> Option<&HeaderValue> {
self.protocol.as_ref()
}
pub fn split(
self,
) -> (
impl Sink<Message, Error = Error> + Send,
impl Stream<Item = Result<Message, Error>> + Send,
) {
let (out_tx, mut out_rx) = tokio::sync::mpsc::channel::<Message>(SPLIT_CHANNEL_CAPACITY);
let (in_tx, in_rx) =
tokio::sync::mpsc::channel::<Result<Message, Error>>(SPLIT_CHANNEL_CAPACITY);
tokio::spawn(async move {
let mut socket = self;
loop {
tokio::select! {
incoming = socket.recv() => {
match incoming {
Some(msg) => {
if in_tx.send(msg).await.is_err() {
break;
}
}
None => break,
}
}
outgoing = out_rx.recv() => {
match outgoing {
Some(msg) => {
if socket.send(msg).await.is_err() {
break;
}
}
None => break,
}
}
}
}
});
(PollSender::new(out_tx).sink_map_err(map_poll_sender_err), SplitStream { rx: in_rx })
}
}
fn map_poll_sender_err<T>(_: tokio_util::sync::PollSendError<T>) -> Error {
Error::Internal("WebSocket connection closed".to_string())
}
struct SplitStream {
rx: tokio::sync::mpsc::Receiver<Result<Message, Error>>,
}
impl Stream for SplitStream {
type Item = Result<Message, Error>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.rx.poll_recv(cx)
}
}