pub const MAX_HEADER_LEN: usize = 14;
pub const OPCODE_CONTINUATION: u8 = 0x0;
pub const OPCODE_TEXT: u8 = 0x1;
pub const OPCODE_BINARY: u8 = 0x2;
pub const OPCODE_CLOSE: u8 = 0x8;
pub const OPCODE_PING: u8 = 0x9;
pub const OPCODE_PONG: u8 = 0xA;
#[derive(Debug, Clone, Copy)]
pub struct WsHeader {
pub fin: bool,
pub opcode: u8,
pub masked: bool,
pub payload_len: u64,
pub mask_key: [u8; 4],
pub header_len: usize,
}
pub fn parse_header(buf: &[u8]) -> Option<WsHeader> {
if buf.len() < 2 {
return None;
}
let fin = buf[0] & 0x80 != 0;
let opcode = buf[0] & 0x0F;
let masked = buf[1] & 0x80 != 0;
let len7 = (buf[1] & 0x7F) as u64;
let (payload_len, mut pos) = if len7 <= 125 {
(len7, 2)
} else if len7 == 126 {
if buf.len() < 4 {
return None;
}
let len = u16::from_be_bytes([buf[2], buf[3]]) as u64;
(len, 4)
} else {
if buf.len() < 10 {
return None;
}
let len = u64::from_be_bytes(buf[2..10].try_into().unwrap());
(len, 10)
};
let mut mask_key = [0u8; 4];
if masked {
if buf.len() < pos + 4 {
return None;
}
mask_key.copy_from_slice(&buf[pos..pos + 4]);
pos += 4;
}
Some(WsHeader {
fin,
opcode,
masked,
payload_len,
mask_key,
header_len: pos,
})
}
pub fn write_binary_header(buf: &mut [u8], payload_len: usize) -> usize {
debug_assert!(buf.len() >= MAX_HEADER_LEN);
buf[0] = 0x82;
if payload_len <= 125 {
buf[1] = payload_len as u8;
2
} else if payload_len <= 65535 {
buf[1] = 126;
buf[2..4].copy_from_slice(&(payload_len as u16).to_be_bytes());
4
} else {
buf[1] = 127;
buf[2..10].copy_from_slice(&(payload_len as u64).to_be_bytes());
10
}
}
pub fn write_masked_binary_header(buf: &mut [u8], payload_len: usize, mask_key: [u8; 4]) -> usize {
let n = write_binary_header(buf, payload_len);
buf[1] |= 0x80;
buf[n..n + 4].copy_from_slice(&mask_key);
n + 4
}
pub fn write_close_frame(buf: &mut [u8]) -> usize {
buf[0] = 0x88;
buf[1] = 0x02; buf[2..4].copy_from_slice(&1000u16.to_be_bytes()); 4
}
pub fn write_masked_close_frame(buf: &mut [u8], mask_key: [u8; 4]) -> usize {
buf[0] = 0x88;
buf[1] = 0x80 | 0x02;
buf[2..6].copy_from_slice(&mask_key);
let status = 1000u16.to_be_bytes();
buf[6] = status[0] ^ mask_key[0];
buf[7] = status[1] ^ mask_key[1];
8
}
pub fn write_masked_pong_frame(buf: &mut [u8], ping_payload: &[u8], mask_key: [u8; 4]) -> usize {
let len = ping_payload.len();
debug_assert!(len <= 125); buf[0] = 0x8A; buf[1] = 0x80 | len as u8;
buf[2..6].copy_from_slice(&mask_key);
for i in 0..len {
buf[6 + i] = ping_payload[i] ^ mask_key[i % 4];
}
6 + len
}
pub fn write_pong_frame(buf: &mut [u8], ping_payload: &[u8]) -> usize {
let len = ping_payload.len();
debug_assert!(len <= 125); buf[0] = 0x8A; buf[1] = len as u8;
if len > 0 {
buf[2..2 + len].copy_from_slice(ping_payload);
}
2 + len
}
#[inline]
pub fn unmask(payload: &mut [u8], mask_key: [u8; 4]) {
let mask_u32 = u32::from_ne_bytes(mask_key);
let mask_u64 = (mask_u32 as u64) | ((mask_u32 as u64) << 32);
let (prefix, aligned, suffix) = unsafe { payload.align_to_mut::<u64>() };
for (i, byte) in prefix.iter_mut().enumerate() {
*byte ^= mask_key[i % 4];
}
let offset = prefix.len() % 4;
let aligned_mask = if offset == 0 {
mask_u64
} else {
let shifted = mask_key.repeat(3);
let start = offset;
let chunk = &shifted[start..start + 8];
u64::from_ne_bytes(chunk.try_into().unwrap())
};
for word in aligned.iter_mut() {
*word ^= aligned_mask;
}
let total_prefix_aligned = prefix.len() + aligned.len() * 8;
for (i, byte) in suffix.iter_mut().enumerate() {
*byte ^= mask_key[(total_prefix_aligned + i) % 4];
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_short_payload() {
let buf = [0x82, 0x05];
let h = parse_header(&buf).unwrap();
assert!(h.fin);
assert_eq!(h.opcode, OPCODE_BINARY);
assert!(!h.masked);
assert_eq!(h.payload_len, 5);
assert_eq!(h.header_len, 2);
}
#[test]
fn test_parse_medium_payload() {
let mut buf = [0u8; 4];
buf[0] = 0x82;
buf[1] = 126;
buf[2..4].copy_from_slice(&300u16.to_be_bytes());
let h = parse_header(&buf).unwrap();
assert_eq!(h.payload_len, 300);
assert_eq!(h.header_len, 4);
}
#[test]
fn test_parse_large_payload() {
let mut buf = [0u8; 10];
buf[0] = 0x82;
buf[1] = 127;
buf[2..10].copy_from_slice(&70000u64.to_be_bytes());
let h = parse_header(&buf).unwrap();
assert_eq!(h.payload_len, 70000);
assert_eq!(h.header_len, 10);
}
#[test]
fn test_parse_masked() {
let buf = [0x82, 0x85, 1, 2, 3, 4];
let h = parse_header(&buf).unwrap();
assert!(h.masked);
assert_eq!(h.payload_len, 5);
assert_eq!(h.mask_key, [1, 2, 3, 4]);
assert_eq!(h.header_len, 6);
}
#[test]
fn test_parse_incomplete() {
assert!(parse_header(&[0x82]).is_none());
assert!(parse_header(&[0x82, 126, 0]).is_none()); assert!(parse_header(&[0x82, 0x85, 1, 2]).is_none()); }
#[test]
fn test_write_binary_header_short() {
let mut buf = [0u8; MAX_HEADER_LEN];
let n = write_binary_header(&mut buf, 50);
assert_eq!(n, 2);
assert_eq!(buf[0], 0x82);
assert_eq!(buf[1], 50);
}
#[test]
fn test_write_binary_header_medium() {
let mut buf = [0u8; MAX_HEADER_LEN];
let n = write_binary_header(&mut buf, 300);
assert_eq!(n, 4);
assert_eq!(buf[1], 126);
assert_eq!(u16::from_be_bytes([buf[2], buf[3]]), 300);
}
#[test]
fn test_write_binary_header_large() {
let mut buf = [0u8; MAX_HEADER_LEN];
let n = write_binary_header(&mut buf, 70000);
assert_eq!(n, 10);
assert_eq!(buf[1], 127);
assert_eq!(u64::from_be_bytes(buf[2..10].try_into().unwrap()), 70000);
}
#[test]
fn test_unmask_roundtrip() {
let mask_key = [0xAB, 0xCD, 0xEF, 0x01];
let original = b"Hello, WebSocket!".to_vec();
let mut data = original.clone();
unmask(&mut data, mask_key);
assert_ne!(data, original);
unmask(&mut data, mask_key); assert_eq!(data, original);
}
#[test]
fn test_unmask_empty() {
let mask_key = [1, 2, 3, 4];
let mut data = vec![];
unmask(&mut data, mask_key); }
#[test]
fn test_close_frame() {
let mut buf = [0u8; 16];
let n = write_close_frame(&mut buf);
assert_eq!(n, 4);
assert_eq!(buf[0], 0x88); assert_eq!(buf[1], 2); assert_eq!(u16::from_be_bytes([buf[2], buf[3]]), 1000);
}
}