use std::io::{self, Read, Write};
use std::net::{Shutdown, TcpStream};
use bytes::{Bytes, BytesMut};
use crate::ws::frame::{FrameError, FrameHeader, Opcode, encode_client_frame};
use crate::ws::mask::MaskKeySource;
use crate::ws::nosigpipe::NoSigpipeTcp;
const INITIAL_RECV_CAPACITY: usize = 4 * 1024 * 1024;
const READ_CHUNK: usize = 4 * 1024 * 1024;
pub(crate) enum Stream {
Plain(NoSigpipeTcp),
Tls(Box<rustls::StreamOwned<rustls::ClientConnection, NoSigpipeTcp>>),
}
impl Stream {
pub(crate) fn tcp_mut(&mut self) -> &mut TcpStream {
match self {
Stream::Plain(s) => s.tcp_mut(),
Stream::Tls(s) => s.sock.tcp_mut(),
}
}
pub(crate) fn shutdown(&mut self) {
let _ = self.tcp_mut().shutdown(Shutdown::Both);
}
}
impl Read for Stream {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
match self {
Stream::Plain(s) => s.read(buf),
Stream::Tls(s) => s.read(buf),
}
}
}
impl Write for Stream {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
match self {
Stream::Plain(s) => s.write(buf),
Stream::Tls(s) => s.write(buf),
}
}
fn flush(&mut self) -> io::Result<()> {
match self {
Stream::Plain(s) => s.flush(),
Stream::Tls(s) => s.flush(),
}
}
}
#[derive(Debug)]
pub(crate) enum WsReadError {
Io(io::Error),
Protocol(String),
ServerClose { code: Option<u16> },
}
impl From<io::Error> for WsReadError {
fn from(e: io::Error) -> Self {
WsReadError::Io(e)
}
}
pub(crate) struct WsClient {
stream: Stream,
recv: BytesMut,
mask_keys: MaskKeySource,
max_payload: usize,
}
impl WsClient {
pub(crate) fn new(
stream: Stream,
leftover: Vec<u8>,
mask_keys: MaskKeySource,
max_payload: usize,
) -> Self {
let mut recv = BytesMut::with_capacity(INITIAL_RECV_CAPACITY.max(leftover.len()));
recv.extend_from_slice(&leftover);
Self {
stream,
recv,
mask_keys,
max_payload,
}
}
pub(crate) fn stream_mut(&mut self) -> &mut Stream {
&mut self.stream
}
pub(crate) fn read_binary_frame(&mut self) -> Result<Bytes, WsReadError> {
loop {
let header = match FrameHeader::parse(&self.recv) {
Ok(h) => h,
Err(FrameError::Incomplete) => {
self.fill_more()?;
continue;
}
Err(FrameError::Protocol(msg)) => {
return Err(WsReadError::Protocol(msg.to_string()));
}
};
let payload_len = header.payload_len as usize;
if payload_len > self.max_payload {
return Err(WsReadError::Protocol(format!(
"WS payload {} bytes exceeds cap {}",
payload_len, self.max_payload
)));
}
let total = header.header_len + payload_len;
if self.recv.len() < total {
self.fill_more()?;
continue;
}
self.recv.advance_to(header.header_len);
let payload = self.recv.split_to(payload_len).freeze();
match header.opcode {
Opcode::Binary => {
if !header.fin {
return Err(WsReadError::Protocol(
"fragmented binary frame; QWP frames are never fragmented".to_string(),
));
}
return Ok(payload);
}
Opcode::Ping => {
self.send_frame(Opcode::Pong, &payload)?;
continue;
}
Opcode::Pong => {
continue;
}
Opcode::Close => {
let code = if payload.len() >= 2 {
Some(u16::from_be_bytes([payload[0], payload[1]]))
} else {
None
};
return Err(WsReadError::ServerClose { code });
}
}
}
}
pub(crate) fn write_binary_frame(&mut self, payload: &[u8]) -> io::Result<()> {
self.send_frame(Opcode::Binary, payload)
}
pub(crate) fn send_close(&mut self, code: u16) -> io::Result<()> {
let bytes = code.to_be_bytes();
self.send_frame(Opcode::Close, &bytes)
}
fn send_frame(&mut self, opcode: Opcode, payload: &[u8]) -> io::Result<()> {
let mut out = Vec::with_capacity(payload.len() + 14);
let mask_key = self
.mask_keys
.next_key()
.map_err(|e| io::Error::other(e.0))?;
encode_client_frame(&mut out, opcode, mask_key, payload);
self.stream.write_all(&out)?;
Ok(())
}
fn fill_more(&mut self) -> Result<(), WsReadError> {
let target = self.recv.len().saturating_add(READ_CHUNK);
if target > self.max_payload.saturating_mul(2) {
return Err(WsReadError::Protocol(format!(
"WS recv buffer would grow to {target} bytes, exceeds {} (2 * max_payload)",
self.max_payload * 2
)));
}
self.recv.reserve(READ_CHUNK);
let spare = self.recv.spare_capacity_mut();
let slice =
unsafe { std::slice::from_raw_parts_mut(spare.as_mut_ptr() as *mut u8, spare.len()) };
let n = self.stream.read(slice)?;
if n == 0 {
return Err(WsReadError::Io(io::Error::new(
io::ErrorKind::UnexpectedEof,
"WS peer closed connection mid-frame",
)));
}
unsafe { self.recv.set_len(self.recv.len() + n) };
Ok(())
}
}
trait AdvanceTo {
fn advance_to(&mut self, n: usize);
}
impl AdvanceTo for BytesMut {
fn advance_to(&mut self, n: usize) {
use bytes::Buf;
Buf::advance(self, n);
}
}
#[cfg(test)]
mod tests {
use crate::ws::frame::Opcode;
fn server_frame(opcode: u8, payload: &[u8]) -> Vec<u8> {
let mut out = Vec::with_capacity(payload.len() + 10);
out.push(0x80 | opcode);
let len = payload.len();
if len <= 125 {
out.push(len as u8);
} else if len <= 0xFFFF {
out.push(126);
out.extend_from_slice(&(len as u16).to_be_bytes());
} else {
out.push(127);
out.extend_from_slice(&(len as u64).to_be_bytes());
}
out.extend_from_slice(payload);
out
}
#[test]
fn server_frame_helper_round_trips() {
let bytes = server_frame(0x02, b"hello");
let header = crate::ws::frame::FrameHeader::parse(&bytes).unwrap();
assert_eq!(header.opcode, Opcode::Binary);
assert_eq!(header.payload_len, 5);
}
}