use crate::courierust_error::{Error, Result};
use alloc::boxed::Box;
use alloc::vec::Vec;
pub const MAX_HEADER_LEN: usize = 14;
pub const MAX_CONTROL_PAYLOAD: usize = 125;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum OpCode {
Continuation,
Text,
Binary,
Close,
Ping,
Pong,
}
impl OpCode {
#[inline]
pub fn from_u8(v: u8) -> Option<Self> {
Some(match v {
0x0 => Self::Continuation,
0x1 => Self::Text,
0x2 => Self::Binary,
0x8 => Self::Close,
0x9 => Self::Ping,
0xa => Self::Pong,
_ => return None,
})
}
#[inline]
pub fn to_u8(self) -> u8 {
match self {
Self::Continuation => 0x0,
Self::Text => 0x1,
Self::Binary => 0x2,
Self::Close => 0x8,
Self::Ping => 0x9,
Self::Pong => 0xa,
}
}
#[inline]
pub fn is_control(self) -> bool {
self.to_u8() & 0x08 != 0
}
#[inline]
pub fn is_data(self) -> bool {
!self.is_control()
}
pub fn as_str(self) -> &'static str {
match self {
Self::Continuation => "continuation",
Self::Text => "text",
Self::Binary => "binary",
Self::Close => "close",
Self::Ping => "ping",
Self::Pong => "pong",
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct FrameHeader {
pub fin: bool,
pub rsv1: bool,
pub rsv2: bool,
pub rsv3: bool,
pub opcode: OpCode,
pub masked: bool,
pub mask_key: [u8; 4],
pub payload_len: u64,
pub header_len: usize,
}
impl FrameHeader {
pub fn data(opcode: OpCode, fin: bool, payload_len: u64) -> Self {
Self {
fin,
rsv1: false,
rsv2: false,
rsv3: false,
opcode,
masked: false,
mask_key: [0; 4],
payload_len,
header_len: 0,
}
}
pub fn parse(buf: &[u8]) -> Result<Option<Self>> {
if buf.len() < 2 {
return Ok(None);
}
let b0 = buf[0];
let b1 = buf[1];
let opcode = OpCode::from_u8(b0 & 0x0f)
.ok_or_else(|| Error::protocol("websocket: reserved opcode received"))?;
let fin = b0 & 0x80 != 0;
let masked = b1 & 0x80 != 0;
let mut need = 2usize;
let len7 = b1 & 0x7f;
let payload_len: u64 = match len7 {
126 => {
need += 2;
if buf.len() < need {
return Ok(None);
}
let v = u16::from_be_bytes([buf[2], buf[3]]) as u64;
if v < 126 {
return Err(Error::protocol(
"websocket: non-minimal 16-bit payload length",
));
}
v
}
127 => {
need += 8;
if buf.len() < need {
return Ok(None);
}
let mut b = [0u8; 8];
b.copy_from_slice(&buf[2..10]);
let v = u64::from_be_bytes(b);
if v >> 63 != 0 {
return Err(Error::protocol("websocket: payload length msb set"));
}
if v < 65_536 {
return Err(Error::protocol(
"websocket: non-minimal 64-bit payload length",
));
}
v
}
n => u64::from(n),
};
let mut mask_key = [0u8; 4];
if masked {
need += 4;
if buf.len() < need {
return Ok(None);
}
mask_key.copy_from_slice(&buf[need - 4..need]);
}
if opcode.is_control() {
if !fin {
return Err(Error::protocol("websocket: fragmented control frame"));
}
if payload_len > MAX_CONTROL_PAYLOAD as u64 {
return Err(Error::protocol(
"websocket: control frame payload too large",
));
}
}
Ok(Some(Self {
fin,
rsv1: b0 & 0x40 != 0,
rsv2: b0 & 0x20 != 0,
rsv3: b0 & 0x10 != 0,
opcode,
masked,
mask_key,
payload_len,
header_len: need,
}))
}
#[inline]
pub fn header_len_hint(buf: &[u8]) -> Option<usize> {
if buf.len() < 2 {
return None;
}
let masked = buf[1] & 0x80 != 0;
let base = match buf[1] & 0x7f {
126 => 4,
127 => 10,
_ => 2,
};
Some(base + if masked { 4 } else { 0 })
}
pub fn check_reserved(&self, allow_rsv1: bool) -> Result<()> {
let ok_rsv1 = self.rsv1 && allow_rsv1 && self.opcode.is_data();
if (self.rsv1 && !ok_rsv1) || self.rsv2 || self.rsv3 {
return Err(Error::protocol(
"websocket: reserved bits set without a negotiated extension",
));
}
Ok(())
}
pub fn write(&self, out: &mut [u8]) -> usize {
debug_assert!(out.len() >= MAX_HEADER_LEN);
let mut b0 = self.opcode.to_u8();
if self.fin {
b0 |= 0x80;
}
if self.rsv1 {
b0 |= 0x40;
}
if self.rsv2 {
b0 |= 0x20;
}
if self.rsv3 {
b0 |= 0x10;
}
out[0] = b0;
let mask_bit = if self.masked { 0x80 } else { 0 };
let mut n = if self.payload_len < 126 {
out[1] = mask_bit | self.payload_len as u8;
2
} else if self.payload_len <= u64::from(u16::MAX) {
out[1] = mask_bit | 126;
out[2..4].copy_from_slice(&(self.payload_len as u16).to_be_bytes());
4
} else {
out[1] = mask_bit | 127;
out[2..10].copy_from_slice(&self.payload_len.to_be_bytes());
10
};
if self.masked {
out[n..n + 4].copy_from_slice(&self.mask_key);
n += 4;
}
n
}
}
#[derive(Clone, Copy, PartialEq, Eq)]
pub struct Mask {
key: [u8; 4],
wide: [u128; 4],
word: [u32; 4],
}
impl core::fmt::Debug for Mask {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.write_str("Mask(<4 bytes>)")
}
}
impl Mask {
pub fn new(key: [u8; 4]) -> Self {
let mut wide = [0u128; 4];
let mut word = [0u32; 4];
for phase in 0..4 {
let mut h = [0u8; 16];
for (j, byte) in h.iter_mut().enumerate() {
*byte = key[(phase + j) & 3];
}
wide[phase] = u128::from_ne_bytes(h);
let mut w = [0u8; 4];
for (j, byte) in w.iter_mut().enumerate() {
*byte = key[(phase + j) & 3];
}
word[phase] = u32::from_ne_bytes(w);
}
Self { key, wide, word }
}
#[inline]
pub fn key(&self) -> [u8; 4] {
self.key
}
pub fn apply(&self, offset: usize, data: &mut [u8]) {
let phase = offset & 3;
let wide = self.wide[phase];
let mut i = 0usize;
while i + 16 <= data.len() {
let lane = u128::from_ne_bytes(
data[i..i + 16]
.try_into()
.expect("slice of exactly 16 bytes"),
) ^ wide;
data[i..i + 16].copy_from_slice(&lane.to_ne_bytes());
i += 16;
}
let word = self.word[phase];
while i + 4 <= data.len() {
let lane =
u32::from_ne_bytes(data[i..i + 4].try_into().expect("slice of exactly 4 bytes"))
^ word;
data[i..i + 4].copy_from_slice(&lane.to_ne_bytes());
i += 4;
}
while i < data.len() {
data[i] ^= self.key[(phase + i) & 3];
i += 1;
}
}
pub fn apply_into(&self, offset: usize, src: &[u8], dst: &mut [u8]) {
debug_assert_eq!(src.len(), dst.len());
let phase = offset & 3;
let wide = self.wide[phase];
let mut i = 0usize;
while i + 16 <= src.len() {
let lane = u128::from_ne_bytes(
src[i..i + 16]
.try_into()
.expect("slice of exactly 16 bytes"),
) ^ wide;
dst[i..i + 16].copy_from_slice(&lane.to_ne_bytes());
i += 16;
}
let word = self.word[phase];
while i + 4 <= src.len() {
let lane =
u32::from_ne_bytes(src[i..i + 4].try_into().expect("slice of exactly 4 bytes"))
^ word;
dst[i..i + 4].copy_from_slice(&lane.to_ne_bytes());
i += 4;
}
while i < src.len() {
dst[i] = src[i] ^ self.key[(phase + i) & 3];
i += 1;
}
}
}
pub trait FrameSink {
fn write_frame(&mut self, header: &[u8], payload: &[u8], mask: Option<Mask>) -> Result<()>;
fn flush(&mut self) -> Result<()> {
Ok(())
}
}
pub const MASK_WINDOW: usize = 64 * 1024;
pub const MASK_WINDOW_MAX: usize = 1024 * 1024;
pub struct StreamSink<W: crate::courierust_io::Write> {
writer: W,
stage: Vec<u8>,
}
impl<W: crate::courierust_io::Write> StreamSink<W> {
pub fn new(writer: W) -> Self {
Self {
writer,
stage: Vec::new(),
}
}
pub fn get_ref(&self) -> &W {
&self.writer
}
pub fn get_mut(&mut self) -> &mut W {
&mut self.writer
}
pub fn into_inner(self) -> W {
self.writer
}
pub fn stage_capacity(&self) -> usize {
self.stage.capacity()
}
}
impl<W: crate::courierust_io::Write> FrameSink for StreamSink<W> {
fn write_frame(&mut self, header: &[u8], payload: &[u8], mask: Option<Mask>) -> Result<()> {
let total = header.len() + payload.len();
if total <= MASK_WINDOW {
self.stage.clear();
self.stage.reserve(total);
self.stage.extend_from_slice(header);
let start = self.stage.len();
self.stage.extend_from_slice(payload);
if let Some(m) = mask {
m.apply(0, &mut self.stage[start..]);
}
let (stage, writer) = (&self.stage, &mut self.writer);
return writer.write_all(stage);
}
self.writer.write_all(header)?;
match mask {
None => self.writer.write_all(payload),
Some(m) => {
let mut off = 0usize;
while off < payload.len() {
let take = core::cmp::min(MASK_WINDOW_MAX, payload.len() - off);
self.stage.clear();
self.stage.extend_from_slice(&payload[off..off + take]);
m.apply(off, &mut self.stage);
let (stage, writer) = (&self.stage, &mut self.writer);
writer.write_all(stage)?;
off += take;
}
Ok(())
}
}
}
fn flush(&mut self) -> Result<()> {
self.writer.flush()
}
}
impl FrameSink for Box<dyn FrameSink + Send> {
fn write_frame(&mut self, header: &[u8], payload: &[u8], mask: Option<Mask>) -> Result<()> {
(**self).write_frame(header, payload, mask)
}
fn flush(&mut self) -> Result<()> {
(**self).flush()
}
}
#[cfg(feature = "std")]
pub struct SharedSink<W: crate::courierust_io::Write> {
inner: std::sync::Arc<std::sync::Mutex<StreamSink<W>>>,
}
#[cfg(feature = "std")]
impl<W: crate::courierust_io::Write> Clone for SharedSink<W> {
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
}
}
}
#[cfg(feature = "std")]
impl<W: crate::courierust_io::Write> SharedSink<W> {
pub fn new(writer: W) -> Self {
Self {
inner: std::sync::Arc::new(std::sync::Mutex::new(StreamSink::new(writer))),
}
}
pub fn with<R>(&self, f: impl FnOnce(&mut W) -> R) -> R {
let mut guard = self
.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
f(guard.get_mut())
}
pub fn transport(&self) -> W
where
W: Clone,
{
self.with(|w| w.clone())
}
}
#[cfg(feature = "std")]
impl<W: crate::courierust_io::Write> FrameSink for SharedSink<W> {
fn write_frame(&mut self, header: &[u8], payload: &[u8], mask: Option<Mask>) -> Result<()> {
let mut guard = self
.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
guard.write_frame(header, payload, mask)
}
fn flush(&mut self) -> Result<()> {
let mut guard = self
.inner
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner());
guard.flush()
}
}
pub mod close {
use crate::courierust_error::{Error, Result};
use crate::courierust_ws::utf8::Utf8Validator;
use alloc::string::String;
use alloc::vec::Vec;
pub const NORMAL: u16 = 1000;
pub const GOING_AWAY: u16 = 1001;
pub const PROTOCOL_ERROR: u16 = 1002;
pub const UNSUPPORTED: u16 = 1003;
pub const NO_STATUS: u16 = 1005;
pub const ABNORMAL: u16 = 1006;
pub const INVALID_PAYLOAD: u16 = 1007;
pub const POLICY: u16 = 1008;
pub const TOO_BIG: u16 = 1009;
pub const MANDATORY_EXTENSION: u16 = 1010;
pub const INTERNAL: u16 = 1011;
pub const SERVICE_RESTART: u16 = 1012;
pub const TRY_AGAIN_LATER: u16 = 1013;
pub const BAD_GATEWAY: u16 = 1014;
pub const TLS_HANDSHAKE: u16 = 1015;
pub fn is_valid_received(code: u16) -> bool {
matches!(code, 1000..=1003 | 1007..=1014 | 3000..=4999)
}
pub fn is_valid_to_send(code: u16) -> bool {
matches!(code, 1000..=1003 | 1007..=1014 | 3000..=4999)
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct CloseFrame {
pub code: u16,
pub reason: String,
}
impl CloseFrame {
pub fn new(code: u16, reason: &str) -> Self {
Self {
code,
reason: String::from(reason),
}
}
pub fn encode(&self) -> Vec<u8> {
let mut out = Vec::with_capacity(2 + self.reason.len());
out.extend_from_slice(&self.code.to_be_bytes());
out.extend_from_slice(self.reason.as_bytes());
out
}
}
pub fn parse(payload: &[u8]) -> Result<Option<CloseFrame>> {
if payload.is_empty() {
return Ok(None);
}
if payload.len() == 1 {
return Err(Error::protocol("websocket: 1-byte close payload"));
}
let code = u16::from_be_bytes([payload[0], payload[1]]);
if !is_valid_received(code) {
return Err(Error::protocol("websocket: invalid close code"));
}
let reason = &payload[2..];
if !Utf8Validator::validate(reason) {
return Err(Error::with_message(
crate::courierust_error::ErrorKind::Protocol,
"websocket: close reason is not UTF-8",
));
}
let text = core::str::from_utf8(reason)
.map_err(|_| Error::protocol("websocket: close reason is not UTF-8"))?;
Ok(Some(CloseFrame {
code,
reason: String::from(text),
}))
}
}
#[cfg(test)]
mod tests {
use super::*;
use alloc::vec;
#[test]
fn small_unmasked_header_roundtrip() {
let h = FrameHeader {
fin: true,
rsv1: false,
rsv2: false,
rsv3: false,
opcode: OpCode::Text,
masked: false,
mask_key: [0; 4],
payload_len: 5,
header_len: 2,
};
let mut buf = [0u8; MAX_HEADER_LEN];
let n = h.write(&mut buf);
assert_eq!(n, 2);
assert_eq!(&buf[..2], &[0x81, 0x05]);
let parsed = FrameHeader::parse(&buf[..n]).unwrap().unwrap();
assert_eq!(parsed.opcode, OpCode::Text);
assert!(parsed.fin);
assert_eq!(parsed.payload_len, 5);
assert_eq!(parsed.header_len, 2);
}
#[test]
fn masked_16_bit_length_roundtrip() {
let h = FrameHeader {
fin: false,
rsv1: true,
rsv2: false,
rsv3: false,
opcode: OpCode::Binary,
masked: true,
mask_key: [0xde, 0xad, 0xbe, 0xef],
payload_len: 4096,
header_len: 0,
};
let mut buf = [0u8; MAX_HEADER_LEN];
let n = h.write(&mut buf);
assert_eq!(n, 8);
let parsed = FrameHeader::parse(&buf[..n]).unwrap().unwrap();
assert_eq!(parsed, FrameHeader { header_len: 8, ..h });
}
#[test]
fn masked_64_bit_length_roundtrip() {
let h = FrameHeader {
fin: true,
rsv1: false,
rsv2: false,
rsv3: false,
opcode: OpCode::Binary,
masked: true,
mask_key: [1, 2, 3, 4],
payload_len: 1 << 40,
header_len: 0,
};
let mut buf = [0u8; MAX_HEADER_LEN];
let n = h.write(&mut buf);
assert_eq!(n, 14);
let parsed = FrameHeader::parse(&buf[..n]).unwrap().unwrap();
assert_eq!(parsed.payload_len, 1 << 40);
assert_eq!(parsed.mask_key, [1, 2, 3, 4]);
}
#[test]
fn partial_headers_ask_for_more_bytes() {
let h = FrameHeader {
fin: true,
rsv1: false,
rsv2: false,
rsv3: false,
opcode: OpCode::Binary,
masked: true,
mask_key: [9, 9, 9, 9],
payload_len: 70_000,
header_len: 0,
};
let mut buf = [0u8; MAX_HEADER_LEN];
let n = h.write(&mut buf);
for cut in 0..n {
assert!(
FrameHeader::parse(&buf[..cut]).unwrap().is_none(),
"cut at {cut} must ask for more"
);
}
assert!(FrameHeader::parse(&buf[..n]).unwrap().is_some());
assert_eq!(FrameHeader::header_len_hint(&buf[..2]), Some(14));
}
#[test]
fn rejects_reserved_opcodes_and_control_violations() {
assert!(FrameHeader::parse(&[0x83, 0x00]).is_err());
assert!(FrameHeader::parse(&[0x09, 0x00]).is_err());
assert!(FrameHeader::parse(&[0x89, 126, 0x00, 126]).is_err());
assert!(FrameHeader::parse(&[0x81, 126, 0x00, 0x05]).is_err());
assert!(FrameHeader::parse(&[0x81, 127, 0, 0, 0, 0, 0, 0, 0, 5]).is_err());
let mut buf = vec![0x82u8, 127];
buf.extend_from_slice(&(1u64 << 63).to_be_bytes());
assert!(FrameHeader::parse(&buf).is_err());
}
#[test]
fn reserved_bit_gate() {
let mut h = FrameHeader::data(OpCode::Text, true, 0);
h.rsv1 = true;
assert!(h.check_reserved(false).is_err());
assert!(h.check_reserved(true).is_ok());
let mut c = FrameHeader::data(OpCode::Ping, true, 0);
c.rsv1 = true;
assert!(c.check_reserved(true).is_err());
let mut d = FrameHeader::data(OpCode::Text, true, 0);
d.rsv3 = true;
assert!(d.check_reserved(true).is_err());
}
#[test]
fn mask_matches_naive_for_all_phases_and_lengths() {
let key = [0x2b, 0x7e, 0x15, 0x16];
let mask = Mask::new(key);
for offset in 0..16usize {
for len in 0..80usize {
let src: Vec<u8> = (0..len).map(|i| (i as u8).wrapping_mul(37)).collect();
let mut got = src.clone();
mask.apply(offset, &mut got);
let want: Vec<u8> = src
.iter()
.enumerate()
.map(|(i, b)| b ^ key[(offset + i) & 3])
.collect();
assert_eq!(got, want, "offset={offset} len={len}");
let mut dst = vec![0u8; len];
mask.apply_into(offset, &src, &mut dst);
assert_eq!(dst, want, "apply_into offset={offset} len={len}");
mask.apply(offset, &mut got);
assert_eq!(got, src);
}
}
}
#[test]
fn close_payload_validation() {
assert_eq!(close::parse(b"").unwrap(), None);
assert!(close::parse(&[0x03]).is_err());
assert_eq!(
close::parse(&[0x03, 0xe8]).unwrap(),
Some(close::CloseFrame::new(1000, ""))
);
assert_eq!(
close::parse(&[0x03, 0xe8, b'b', b'y', b'e']).unwrap(),
Some(close::CloseFrame::new(1000, "bye"))
);
for code in [0u16, 999, 1004, 1005, 1006, 1015, 1016, 2999, 5000] {
let mut p = code.to_be_bytes().to_vec();
p.push(b'x');
assert!(close::parse(&p).is_err(), "code {code} must be rejected");
}
for code in [
1000u16, 1001, 1002, 1003, 1007, 1008, 1009, 1010, 1011, 1012, 1013, 1014, 3000, 3999,
4000, 4999,
] {
let p = code.to_be_bytes().to_vec();
assert!(close::parse(&p).is_ok(), "code {code} must be accepted");
}
assert!(close::parse(&[0x03, 0xe8, 0xff, 0xfe]).is_err());
assert!(close::parse(&[0x03, 0xe8, 0xe2, 0x82]).is_err());
}
}