#![cfg_attr(not(feature = "_egress"), allow(dead_code))]
use super::mask::apply_mask;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum Opcode {
Binary = 0x2,
Close = 0x8,
Ping = 0x9,
Pong = 0xA,
}
impl Opcode {
fn as_u8(self) -> u8 {
self as u8
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) struct FrameHeader {
pub fin: bool,
pub opcode: Opcode,
pub payload_len: u64,
pub header_len: usize,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum FrameError {
Incomplete,
Protocol(&'static str),
}
const FIN_BIT: u8 = 0x80;
const RSV_BITS: u8 = 0x70;
const OPCODE_MASK: u8 = 0x0F;
const MASK_BIT: u8 = 0x80;
const LEN_MASK: u8 = 0x7F;
pub(crate) const OPCODE_CONTINUATION: u8 = 0x0;
pub(crate) const OPCODE_TEXT: u8 = 0x1;
pub(crate) const OPCODE_BINARY: u8 = 0x2;
pub(crate) const OPCODE_CLOSE: u8 = 0x8;
pub(crate) const OPCODE_PING: u8 = 0x9;
pub(crate) const OPCODE_PONG: u8 = 0xA;
impl FrameHeader {
pub(crate) fn parse(bytes: &[u8]) -> Result<Self, FrameError> {
if bytes.len() < 2 {
return Err(FrameError::Incomplete);
}
let b0 = bytes[0];
let b1 = bytes[1];
if b0 & RSV_BITS != 0 {
return Err(FrameError::Protocol("WS frame has reserved bits set"));
}
let fin = b0 & FIN_BIT != 0;
let opcode = match b0 & OPCODE_MASK {
OPCODE_BINARY => Opcode::Binary,
OPCODE_CLOSE => Opcode::Close,
OPCODE_PING => Opcode::Ping,
OPCODE_PONG => Opcode::Pong,
OPCODE_CONTINUATION => {
return Err(FrameError::Protocol(
"WS continuation frame from server (QWP never fragments)",
));
}
OPCODE_TEXT => {
return Err(FrameError::Protocol("WS text frame (QWP is binary-only)"));
}
_ => {
return Err(FrameError::Protocol("WS frame has reserved opcode"));
}
};
let is_control = matches!(opcode, Opcode::Close | Opcode::Ping | Opcode::Pong);
if is_control && !fin {
return Err(FrameError::Protocol("fragmented control frame"));
}
if b1 & MASK_BIT != 0 {
return Err(FrameError::Protocol("masked frame from server"));
}
let len_field = b1 & LEN_MASK;
let (payload_len, header_len) = match len_field {
0..=125 => (len_field as u64, 2),
126 => {
if bytes.len() < 4 {
return Err(FrameError::Incomplete);
}
let l = u16::from_be_bytes([bytes[2], bytes[3]]) as u64;
if l < 126 {
return Err(FrameError::Protocol(
"16-bit WS length < 126 (must use 7-bit form)",
));
}
(l, 4)
}
127 => {
if bytes.len() < 10 {
return Err(FrameError::Incomplete);
}
let l = u64::from_be_bytes([
bytes[2], bytes[3], bytes[4], bytes[5], bytes[6], bytes[7], bytes[8], bytes[9],
]);
if l >> 63 != 0 {
return Err(FrameError::Protocol("64-bit WS length has high bit set"));
}
if l <= 0xFFFF {
return Err(FrameError::Protocol(
"64-bit WS length ≤ 0xFFFF (must use 16-bit form)",
));
}
(l, 10)
}
_ => unreachable!("len_field is 7 bits"),
};
if is_control && payload_len > 125 {
return Err(FrameError::Protocol("control frame payload > 125 bytes"));
}
Ok(FrameHeader {
fin,
opcode,
payload_len,
header_len,
})
}
}
pub(crate) fn encode_client_frame<'a>(
out: &'a mut Vec<u8>,
opcode: Opcode,
mask_key: [u8; 4],
payload: &[u8],
) -> &'a [u8] {
let start = out.len();
out.push(FIN_BIT | opcode.as_u8());
let len = payload.len();
if len <= 125 {
out.push(MASK_BIT | (len as u8));
} else if len <= 0xFFFF {
out.push(MASK_BIT | 126);
out.extend_from_slice(&(len as u16).to_be_bytes());
} else {
out.push(MASK_BIT | 127);
out.extend_from_slice(&(len as u64).to_be_bytes());
}
out.extend_from_slice(&mask_key);
let payload_start = out.len();
out.extend_from_slice(payload);
apply_mask(&mut out[payload_start..], mask_key, 0);
&out[start..]
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_binary_short() {
let bytes = [0x82, 0x05, 0, 0, 0, 0, 0];
let h = FrameHeader::parse(&bytes).unwrap();
assert!(h.fin);
assert_eq!(h.opcode, Opcode::Binary);
assert_eq!(h.payload_len, 5);
assert_eq!(h.header_len, 2);
}
#[test]
fn parse_binary_16bit_length() {
let bytes = [0x82, 126, 0x03, 0xE8];
let h = FrameHeader::parse(&bytes).unwrap();
assert_eq!(h.payload_len, 1000);
assert_eq!(h.header_len, 4);
}
#[test]
fn parse_binary_64bit_length() {
let bytes = [0x82, 127, 0, 0, 0, 0, 0, 0x10, 0, 0];
let h = FrameHeader::parse(&bytes).unwrap();
assert_eq!(h.payload_len, 0x10_0000);
assert_eq!(h.header_len, 10);
}
#[test]
fn parse_incomplete_returns_incomplete() {
assert_eq!(
FrameHeader::parse(&[0x82]).unwrap_err(),
FrameError::Incomplete
);
assert_eq!(
FrameHeader::parse(&[0x82, 126, 0]).unwrap_err(),
FrameError::Incomplete
);
assert_eq!(
FrameHeader::parse(&[0x82, 127, 0, 0, 0, 0]).unwrap_err(),
FrameError::Incomplete
);
}
#[test]
fn parse_rejects_reserved_bits() {
let bytes = [0xC2, 0x05];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_mask_from_server() {
let bytes = [0x82, 0x80 | 0x05, 0, 0, 0, 0, 1, 2, 3, 4];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_text() {
let bytes = [0x81, 0x05];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_continuation() {
let bytes = [0x80, 0x05];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_reserved_opcode() {
let bytes = [0x8B, 0x05];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_non_minimal_16bit_length() {
let bytes = [0x82, 126, 0, 100];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_non_minimal_64bit_length() {
let bytes = [0x82, 127, 0, 0, 0, 0, 0, 0, 0x03, 0xE8];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_64bit_high_bit() {
let bytes = [0x82, 127, 0x80, 0, 0, 0, 0, 0, 0, 0];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_fragmented_control() {
let bytes = [0x09, 0x00];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_rejects_oversized_control() {
let bytes = [0x89, 126, 0, 200];
assert!(matches!(
FrameHeader::parse(&bytes),
Err(FrameError::Protocol(_))
));
}
#[test]
fn parse_close_and_ping() {
let close = FrameHeader::parse(&[0x88, 0x02, 0x03, 0xE8]).unwrap();
assert_eq!(close.opcode, Opcode::Close);
let ping = FrameHeader::parse(&[0x89, 0x00]).unwrap();
assert_eq!(ping.opcode, Opcode::Ping);
let pong = FrameHeader::parse(&[0x8A, 0x00]).unwrap();
assert_eq!(pong.opcode, Opcode::Pong);
}
#[test]
fn encode_small_binary_frame() {
let mut out = Vec::new();
let payload = b"hello";
let mask = [0x11, 0x22, 0x33, 0x44];
let frame = encode_client_frame(&mut out, Opcode::Binary, mask, payload);
assert_eq!(frame[0], 0x82);
assert_eq!(frame[1], 0x85);
assert_eq!(&frame[2..6], &mask);
let mut payload_check = frame[6..].to_vec();
apply_mask(&mut payload_check, mask, 0);
assert_eq!(payload_check, payload);
}
#[test]
fn encode_medium_frame_uses_16bit_length() {
let mut out = Vec::new();
let payload = vec![0xAB; 1000];
let mask = [0, 0, 0, 0]; let frame = encode_client_frame(&mut out, Opcode::Binary, mask, &payload);
assert_eq!(frame[0], 0x82);
assert_eq!(frame[1], 0x80 | 126);
assert_eq!(u16::from_be_bytes([frame[2], frame[3]]), 1000);
assert_eq!(&frame[8..], &payload[..]);
}
#[test]
fn encode_large_frame_uses_64bit_length() {
let mut out = Vec::new();
let payload = vec![0u8; 0x1_0000]; let mask = [0, 0, 0, 0];
let frame = encode_client_frame(&mut out, Opcode::Binary, mask, &payload);
assert_eq!(frame[0], 0x82);
assert_eq!(frame[1], 0x80 | 127);
assert_eq!(
u64::from_be_bytes([
frame[2], frame[3], frame[4], frame[5], frame[6], frame[7], frame[8], frame[9]
]),
0x1_0000
);
}
#[test]
fn encode_close_frame_zero_payload() {
let mut out = Vec::new();
let frame = encode_client_frame(&mut out, Opcode::Close, [1, 2, 3, 4], b"");
assert_eq!(frame[0], 0x88);
assert_eq!(frame[1], 0x80);
assert_eq!(frame.len(), 6);
assert_eq!(&frame[2..6], &[1, 2, 3, 4]);
}
#[test]
fn round_trip_parser_against_writer() {
let payload = vec![0u8; 200_000];
let mut out = Vec::new();
let mask = [0xAA, 0xBB, 0xCC, 0xDD];
encode_client_frame(&mut out, Opcode::Binary, mask, &payload);
let mut server_view = out.clone();
server_view[1] &= !MASK_BIT;
server_view.drain(10..14);
let header = FrameHeader::parse(&server_view).unwrap();
assert_eq!(header.opcode, Opcode::Binary);
assert_eq!(header.payload_len as usize, payload.len());
}
}