From 5b9bc69d16de24dcb4f20179430e8e65d60e55ad Mon Sep 17 00:00:00 2001 From: Pierre Le Fevre Date: Sun, 17 May 2026 11:39:52 +0200 Subject: [PATCH] Implement WebSocket wire protocol --- Cargo.lock | 1 + crates/crypto/src/lib.rs | 3 +- crates/crypto/src/sha1.rs | 181 ++++++ crates/net/Cargo.toml | 1 + crates/net/src/websocket.rs | 1145 +++++++++++++++++++++++++++-------- 5 files changed, 1094 insertions(+), 237 deletions(-) create mode 100644 crates/crypto/src/sha1.rs diff --git a/Cargo.lock b/Cargo.lock index 7fbabd7..9993459 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -111,6 +111,7 @@ name = "we-net" version = "0.1.0" dependencies = [ "we-crypto", + "we-image", "we-url", ] diff --git a/crates/crypto/src/lib.rs b/crates/crypto/src/lib.rs index e7d732c..8c1921f 100644 --- a/crates/crypto/src/lib.rs +++ b/crates/crypto/src/lib.rs @@ -1,4 +1,4 @@ -//! Pure Rust cryptography — AES-GCM, ChaCha20-Poly1305, SHA-2, X25519, RSA, ECDSA, X.509, ASN.1. +//! Pure Rust cryptography — AES-GCM, ChaCha20-Poly1305, SHA-1, SHA-2, X25519, RSA, ECDSA, X.509, ASN.1. pub mod aes_gcm; pub mod asn1; @@ -8,6 +8,7 @@ pub mod ecdsa; pub mod hkdf; pub mod hmac; pub mod rsa; +pub mod sha1; pub mod sha2; pub mod x25519; pub mod x509; diff --git a/crates/crypto/src/sha1.rs b/crates/crypto/src/sha1.rs new file mode 100644 index 0000000..e959e56 --- /dev/null +++ b/crates/crypto/src/sha1.rs @@ -0,0 +1,181 @@ +//! SHA-1 hash function (FIPS 180-4). +//! +//! SHA-1 is cryptographically broken and must not be used for signatures, +//! passwords, MACs, or new security protocols. It is provided only because +//! the WebSocket opening handshake in RFC 6455 requires +//! `base64(SHA-1(Sec-WebSocket-Key || GUID))`, where the hashed value is +//! public and not used as a secret. + +/// SHA-1 hasher with a streaming API. +#[derive(Clone)] +pub struct Sha1 { + state: [u32; 5], + buf: [u8; 64], + buf_len: usize, + total_len: u64, +} + +impl Sha1 { + pub fn new() -> Self { + Self { + state: [ + 0x6745_2301, + 0xEFCD_AB89, + 0x98BA_DCFE, + 0x1032_5476, + 0xC3D2_E1F0, + ], + buf: [0; 64], + buf_len: 0, + total_len: 0, + } + } + + pub fn update(&mut self, data: &[u8]) { + self.total_len += data.len() as u64; + let mut offset = 0; + + if self.buf_len > 0 { + let copy_len = (64 - self.buf_len).min(data.len()); + self.buf[self.buf_len..self.buf_len + copy_len].copy_from_slice(&data[..copy_len]); + self.buf_len += copy_len; + offset += copy_len; + + if self.buf_len == 64 { + let block = self.buf; + sha1_compress(&mut self.state, &block); + self.buf_len = 0; + } + } + + while offset + 64 <= data.len() { + let block: [u8; 64] = data[offset..offset + 64].try_into().unwrap(); + sha1_compress(&mut self.state, &block); + offset += 64; + } + + let remaining = data.len() - offset; + if remaining > 0 { + self.buf[..remaining].copy_from_slice(&data[offset..]); + self.buf_len = remaining; + } + } + + pub fn finalize(mut self) -> [u8; 20] { + let bit_len = self.total_len * 8; + self.buf[self.buf_len] = 0x80; + self.buf_len += 1; + + if self.buf_len > 56 { + for b in &mut self.buf[self.buf_len..] { + *b = 0; + } + let block = self.buf; + sha1_compress(&mut self.state, &block); + self.buf_len = 0; + } + + for b in &mut self.buf[self.buf_len..56] { + *b = 0; + } + self.buf[56..].copy_from_slice(&bit_len.to_be_bytes()); + let block = self.buf; + sha1_compress(&mut self.state, &block); + + let mut out = [0u8; 20]; + for (i, word) in self.state.iter().enumerate() { + out[i * 4..(i + 1) * 4].copy_from_slice(&word.to_be_bytes()); + } + out + } +} + +impl Default for Sha1 { + fn default() -> Self { + Self::new() + } +} + +/// One-shot SHA-1. +pub fn sha1(data: &[u8]) -> [u8; 20] { + let mut h = Sha1::new(); + h.update(data); + h.finalize() +} + +fn sha1_compress(state: &mut [u32; 5], block: &[u8; 64]) { + let mut w = [0u32; 80]; + for (i, word) in w.iter_mut().take(16).enumerate() { + let j = i * 4; + *word = u32::from_be_bytes([block[j], block[j + 1], block[j + 2], block[j + 3]]); + } + for i in 16..80 { + w[i] = (w[i - 3] ^ w[i - 8] ^ w[i - 14] ^ w[i - 16]).rotate_left(1); + } + + let [mut a, mut b, mut c, mut d, mut e] = *state; + + for (i, &wi) in w.iter().enumerate() { + let (f, k) = match i { + 0..=19 => ((b & c) | ((!b) & d), 0x5A82_7999), + 20..=39 => (b ^ c ^ d, 0x6ED9_EBA1), + 40..=59 => ((b & c) | (b & d) | (c & d), 0x8F1B_BCDC), + _ => (b ^ c ^ d, 0xCA62_C1D6), + }; + let temp = a + .rotate_left(5) + .wrapping_add(f) + .wrapping_add(e) + .wrapping_add(k) + .wrapping_add(wi); + e = d; + d = c; + c = b.rotate_left(30); + b = a; + a = temp; + } + + state[0] = state[0].wrapping_add(a); + state[1] = state[1].wrapping_add(b); + state[2] = state[2].wrapping_add(c); + state[3] = state[3].wrapping_add(d); + state[4] = state[4].wrapping_add(e); +} + +#[cfg(test)] +mod tests { + use super::*; + + fn hex(bytes: &[u8]) -> String { + let mut s = String::new(); + for b in bytes { + s.push_str(&format!("{b:02x}")); + } + s + } + + #[test] + fn sha1_known_vectors() { + assert_eq!(hex(&sha1(b"")), "da39a3ee5e6b4b0d3255bfef95601890afd80709"); + assert_eq!( + hex(&sha1(b"abc")), + "a9993e364706816aba3e25717850c26c9cd0d89d" + ); + assert_eq!( + hex(&sha1(b"The quick brown fox jumps over the lazy dog")), + "2fd4e1c67a2d28fced849ee1bb76e7391b93eb12" + ); + } + + #[test] + fn sha1_streaming_matches_one_shot() { + let mut h = Sha1::new(); + h.update(b"The quick "); + h.update(b"brown fox "); + h.update(b"jumps over the lazy dog"); + assert_eq!( + h.finalize(), + sha1(b"The quick brown fox jumps over the lazy dog") + ); + } +} diff --git a/crates/net/Cargo.toml b/crates/net/Cargo.toml index 721376f..e6e4f5a 100644 --- a/crates/net/Cargo.toml +++ b/crates/net/Cargo.toml @@ -10,3 +10,4 @@ path = "src/lib.rs" [dependencies] we-url = { path = "../url" } we-crypto = { path = "../crypto" } +we-image = { path = "../image" } diff --git a/crates/net/src/websocket.rs b/crates/net/src/websocket.rs index bf6a528..f6e7319 100644 --- a/crates/net/src/websocket.rs +++ b/crates/net/src/websocket.rs @@ -4,6 +4,8 @@ use std::fmt; use std::io::{self, Read, Write}; use std::time::{Duration, SystemTime, UNIX_EPOCH}; +use we_crypto::sha1; +use we_image::deflate; use we_url::Url; use crate::tcp::TcpConnection; @@ -12,6 +14,9 @@ use crate::tls::handshake::{self, HandshakeError, TlsStream}; const GUID: &str = "258EAFA5-E914-47DA-95CA-C5AB0DC85B11"; const DEFAULT_CONNECT_TIMEOUT: Duration = Duration::from_secs(30); const DEFAULT_READ_TIMEOUT: Duration = Duration::from_secs(30); +const DEFAULT_MAX_FRAME_SIZE: usize = 16 * 1024 * 1024; +const DEFAULT_MAX_MESSAGE_SIZE: usize = 16 * 1024 * 1024; +const MAX_CONTROL_PAYLOAD: usize = 125; #[derive(Debug)] pub enum WebSocketError { @@ -22,6 +27,7 @@ pub enum WebSocketError { Io(io::Error), Handshake(String), Protocol(String), + Compression(String), MessageTooLarge, } @@ -35,6 +41,7 @@ impl fmt::Display for WebSocketError { Self::Io(e) => write!(f, "I/O error: {e}"), Self::Handshake(s) => write!(f, "WebSocket handshake failed: {s}"), Self::Protocol(s) => write!(f, "WebSocket protocol error: {s}"), + Self::Compression(s) => write!(f, "WebSocket compression error: {s}"), Self::MessageTooLarge => write!(f, "WebSocket message too large"), } } @@ -58,6 +65,12 @@ impl From for WebSocketError { } } +impl From for WebSocketError { + fn from(err: deflate::DeflateError) -> Self { + Self::Compression(err.to_string()) + } +} + pub type Result = std::result::Result; impl WebSocketError { @@ -87,6 +100,63 @@ pub struct HandshakeResult { pub extensions: String, } +#[derive(Debug, Clone)] +pub struct WebSocketConfig { + pub max_frame_size: usize, + pub max_message_size: usize, + pub read_timeout: Duration, + pub permessage_deflate: Option, +} + +impl Default for WebSocketConfig { + fn default() -> Self { + Self { + max_frame_size: DEFAULT_MAX_FRAME_SIZE, + max_message_size: DEFAULT_MAX_MESSAGE_SIZE, + read_timeout: DEFAULT_READ_TIMEOUT, + permessage_deflate: None, + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct PermessageDeflateConfig { + pub client_no_context_takeover: bool, + pub server_no_context_takeover: bool, + pub client_max_window_bits: Option, + pub server_max_window_bits: Option, +} + +impl Default for PermessageDeflateConfig { + fn default() -> Self { + Self { + client_no_context_takeover: true, + server_no_context_takeover: true, + client_max_window_bits: None, + server_max_window_bits: None, + } + } +} + +impl PermessageDeflateConfig { + fn offer_header(&self) -> String { + let mut parts = vec!["permessage-deflate".to_string()]; + if self.client_no_context_takeover { + parts.push("client_no_context_takeover".to_string()); + } + if self.server_no_context_takeover { + parts.push("server_no_context_takeover".to_string()); + } + if let Some(bits) = self.client_max_window_bits { + parts.push(format!("client_max_window_bits={bits}")); + } + if let Some(bits) = self.server_max_window_bits { + parts.push(format!("server_max_window_bits={bits}")); + } + parts.join("; ") + } +} + enum WsStream { Plain(TcpConnection), Tls(TlsStream), @@ -134,10 +204,23 @@ impl Write for WsStream { pub struct WebSocketClient { stream: WsStream, + config: WebSocketConfig, + negotiated_deflate: Option, + fragments: Option, + sent_close: bool, + received_close: bool, } impl WebSocketClient { pub fn connect(url: &Url, protocols: &[String]) -> Result<(Self, HandshakeResult)> { + Self::connect_with_config(url, protocols, WebSocketConfig::default()) + } + + pub fn connect_with_config( + url: &Url, + protocols: &[String], + config: WebSocketConfig, + ) -> Result<(Self, HandshakeResult)> { let scheme = url.scheme(); let is_tls = match scheme { "ws" => false, @@ -152,7 +235,7 @@ impl WebSocketClient { .ok_or_else(|| WebSocketError::InvalidUrl("missing port".to_string()))?; let tcp = TcpConnection::connect_timeout(&host, port, DEFAULT_CONNECT_TIMEOUT)?; - tcp.set_read_timeout(Some(DEFAULT_READ_TIMEOUT))?; + tcp.set_read_timeout(Some(config.read_timeout))?; let stream = if is_tls { let (tls, _) = handshake::connect_with_alpn(tcp, &host, &["http/1.1"])?; WsStream::Tls(tls) @@ -160,184 +243,506 @@ impl WebSocketClient { WsStream::Plain(tcp) }; - let mut client = Self { stream }; - let key_bytes = random_bytes_16(); - let key = base64_encode(&key_bytes); - let path = request_target(url); - let host_header = if url.port().is_some() { - format!("{host}:{port}") - } else { - host.clone() + let mut client = Self { + stream, + config, + negotiated_deflate: None, + fragments: None, + sent_close: false, + received_close: false, }; - let mut request = format!( - "GET {path} HTTP/1.1\r\n\ - Host: {host_header}\r\n\ - Upgrade: websocket\r\n\ - Connection: Upgrade\r\n\ - Sec-WebSocket-Key: {key}\r\n\ - Sec-WebSocket-Version: 13\r\n" - ); - if !protocols.is_empty() { - request.push_str("Sec-WebSocket-Protocol: "); - request.push_str(&protocols.join(", ")); - request.push_str("\r\n"); - } - request.push_str("\r\n"); + let key = base64_encode(&random_bytes_16()); + let request = build_handshake_request(url, protocols, &client.config, &key)?; client.stream.write_all(request.as_bytes())?; client.stream.flush()?; let (status, headers) = read_http_upgrade_response(&mut client.stream)?; - if status != 101 { - return Err(WebSocketError::Handshake(format!( - "expected status 101, got {status}" - ))); - } - if !header_token_contains(&headers, "upgrade", "websocket") { - return Err(WebSocketError::Handshake("missing Upgrade header".into())); - } - if !header_token_contains(&headers, "connection", "upgrade") { - return Err(WebSocketError::Handshake( - "missing Connection header".into(), - )); - } - let expected_accept = base64_encode(&sha1(format!("{key}{GUID}").as_bytes())); - let accept = header_value(&headers, "sec-websocket-accept").unwrap_or_default(); - if accept.trim() != expected_accept { - return Err(WebSocketError::Handshake( - "invalid Sec-WebSocket-Accept".into(), - )); - } + validate_handshake_response(status, &headers, &key)?; + let extensions = header_value(&headers, "sec-websocket-extensions").unwrap_or_default(); + client.negotiated_deflate = negotiate_permessage_deflate( + client.config.permessage_deflate.as_ref(), + extensions.as_str(), + )?; Ok(( client, HandshakeResult { protocol: header_value(&headers, "sec-websocket-protocol").unwrap_or_default(), - extensions: header_value(&headers, "sec-websocket-extensions").unwrap_or_default(), + extensions, }, )) } - pub fn send_text(&mut self, text: &str) -> Result<()> { - self.write_frame(0x1, text.as_bytes()) - } - pub fn set_read_timeout(&self, duration: Option) -> Result<()> { self.stream.set_read_timeout(duration) } + pub fn send_text(&mut self, text: &str) -> Result<()> { + self.write_message_frame(Opcode::Text, text.as_bytes()) + } + pub fn send_binary(&mut self, bytes: &[u8]) -> Result<()> { - self.write_frame(0x2, bytes) + self.write_message_frame(Opcode::Binary, bytes) } - pub fn close(&mut self, code: Option, reason: &str) -> Result<()> { - let mut payload = Vec::new(); - if let Some(code) = code { - payload.extend_from_slice(&code.to_be_bytes()); - payload.extend_from_slice(reason.as_bytes()); + pub fn send_ping(&mut self, payload: &[u8]) -> Result<()> { + if payload.len() > MAX_CONTROL_PAYLOAD { + return Err(WebSocketError::Protocol( + "control frame payload too large".into(), + )); } - self.write_frame(0x8, &payload) + self.write_frame(Opcode::Ping, payload, true, false) + } + + pub fn send_pong(&mut self, payload: &[u8]) -> Result<()> { + if payload.len() > MAX_CONTROL_PAYLOAD { + return Err(WebSocketError::Protocol( + "control frame payload too large".into(), + )); + } + self.write_frame(Opcode::Pong, payload, true, false) + } + + pub fn close(&mut self, code: Option, reason: &str) -> Result<()> { + let payload = close_payload(code, reason)?; + self.sent_close = true; + self.write_frame(Opcode::Close, &payload, true, false) } pub fn read_message(&mut self) -> Result { - let frame = self.read_frame()?; - match frame.opcode { - 0x1 => { - let text = String::from_utf8(frame.payload) - .map_err(|_| WebSocketError::Protocol("invalid UTF-8 text".into()))?; - Ok(Message::Text(text)) + loop { + let frame = self.read_frame()?; + if frame.is_control() { + return self.handle_control_frame(frame); + } + + match self.handle_data_frame(frame)? { + Some(message) => return Ok(message), + None => continue, } - 0x2 => Ok(Message::Binary(frame.payload)), - 0x8 => { + } + } + + fn write_message_frame(&mut self, opcode: Opcode, payload: &[u8]) -> Result<()> { + if payload.len() > self.config.max_message_size { + return Err(WebSocketError::MessageTooLarge); + } + if self.negotiated_deflate.is_some() { + let compressed = compress_permessage_deflate(payload); + self.write_frame(opcode, &compressed, true, true) + } else { + self.write_frame(opcode, payload, true, false) + } + } + + fn write_frame(&mut self, opcode: Opcode, payload: &[u8], fin: bool, rsv1: bool) -> Result<()> { + let frame = Frame { + fin, + rsv1, + rsv2: false, + rsv3: false, + opcode, + masked: true, + payload: payload.to_vec(), + }; + frame.validate(self.config.max_frame_size, true)?; + frame.write_to(&mut self.stream, Some(random_bytes_4()))?; + self.stream.flush()?; + Ok(()) + } + + fn read_frame(&mut self) -> Result { + Frame::read_from(&mut self.stream, self.config.max_frame_size, false) + } + + fn handle_control_frame(&mut self, frame: Frame) -> Result { + match frame.opcode { + Opcode::Close => { let (code, reason) = parse_close_payload(&frame.payload)?; + self.received_close = true; + if !self.sent_close { + self.sent_close = true; + self.write_frame(Opcode::Close, &frame.payload, true, false)?; + } Ok(Message::Close { code, reason }) } - 0x9 => { - self.write_frame(0xA, &frame.payload)?; + Opcode::Ping => { + self.write_frame(Opcode::Pong, &frame.payload, true, false)?; Ok(Message::Ping(frame.payload)) } - 0xA => Ok(Message::Pong(frame.payload)), - _ => Err(WebSocketError::Protocol("unsupported opcode".into())), + Opcode::Pong => Ok(Message::Pong(frame.payload)), + _ => Err(WebSocketError::Protocol("not a control frame".into())), } } - fn write_frame(&mut self, opcode: u8, payload: &[u8]) -> Result<()> { - if payload.len() > u32::MAX as usize { - return Err(WebSocketError::MessageTooLarge); - } - let mut frame = Vec::with_capacity(payload.len() + 14); - frame.push(0x80 | (opcode & 0x0F)); - let mask_bit = 0x80; - match payload.len() { - 0..=125 => frame.push(mask_bit | payload.len() as u8), - 126..=65535 => { - frame.push(mask_bit | 126); - frame.extend_from_slice(&(payload.len() as u16).to_be_bytes()); + fn handle_data_frame(&mut self, frame: Frame) -> Result> { + match frame.opcode { + Opcode::Text | Opcode::Binary => { + if self.fragments.is_some() { + return Err(WebSocketError::Protocol( + "new data frame before fragmented message completed".into(), + )); + } + if frame.fin { + return self.message_from_parts(frame.opcode, frame.rsv1, frame.payload); + } + self.fragments = Some(FragmentedMessage { + opcode: frame.opcode, + compressed: frame.rsv1, + payload: frame.payload, + }); + Ok(None) } - _ => { - frame.push(mask_bit | 127); - frame.extend_from_slice(&(payload.len() as u64).to_be_bytes()); + Opcode::Continuation => { + let mut fragments = self.fragments.take().ok_or_else(|| { + WebSocketError::Protocol("unexpected continuation frame".into()) + })?; + fragments.payload.extend_from_slice(&frame.payload); + if fragments.payload.len() > self.config.max_message_size { + return Err(WebSocketError::MessageTooLarge); + } + if frame.fin { + self.message_from_parts( + fragments.opcode, + fragments.compressed, + fragments.payload, + ) + } else { + self.fragments = Some(fragments); + Ok(None) + } } + _ => Err(WebSocketError::Protocol("unsupported data opcode".into())), } - let mask = random_bytes_4(); - frame.extend_from_slice(&mask); - for (i, b) in payload.iter().enumerate() { - frame.push(*b ^ mask[i % 4]); + } + + fn message_from_parts( + &mut self, + opcode: Opcode, + compressed: bool, + payload: Vec, + ) -> Result> { + if payload.len() > self.config.max_message_size { + return Err(WebSocketError::MessageTooLarge); + } + let payload = if compressed { + decompress_permessage_deflate(&payload)? + } else { + payload + }; + if payload.len() > self.config.max_message_size { + return Err(WebSocketError::MessageTooLarge); + } + match opcode { + Opcode::Text => match String::from_utf8(payload) { + Ok(text) => Ok(Some(Message::Text(text))), + Err(_) => { + self.close(Some(1007), "")?; + Err(WebSocketError::Protocol("invalid UTF-8 text".into())) + } + }, + Opcode::Binary => Ok(Some(Message::Binary(payload))), + _ => Err(WebSocketError::Protocol("invalid message opcode".into())), } - self.stream.write_all(&frame)?; - self.stream.flush()?; - Ok(()) } - fn read_frame(&mut self) -> Result { + pub fn close_handshake_complete(&self) -> bool { + self.sent_close && self.received_close + } +} + +struct FragmentedMessage { + opcode: Opcode, + compressed: bool, + payload: Vec, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum Opcode { + Continuation, + Text, + Binary, + Close, + Ping, + Pong, +} + +impl Opcode { + fn from_wire(value: u8) -> Result { + match value { + 0x0 => Ok(Self::Continuation), + 0x1 => Ok(Self::Text), + 0x2 => Ok(Self::Binary), + 0x8 => Ok(Self::Close), + 0x9 => Ok(Self::Ping), + 0xA => Ok(Self::Pong), + _ => Err(WebSocketError::Protocol(format!( + "unsupported opcode 0x{value:x}" + ))), + } + } + + fn wire(self) -> u8 { + match self { + Self::Continuation => 0x0, + Self::Text => 0x1, + Self::Binary => 0x2, + Self::Close => 0x8, + Self::Ping => 0x9, + Self::Pong => 0xA, + } + } + + fn is_control(self) -> bool { + matches!(self, Self::Close | Self::Ping | Self::Pong) + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct Frame { + pub fin: bool, + pub rsv1: bool, + pub rsv2: bool, + pub rsv3: bool, + pub opcode: Opcode, + pub masked: bool, + pub payload: Vec, +} + +impl Frame { + pub fn is_control(&self) -> bool { + self.opcode.is_control() + } + + pub fn read_from( + read: &mut R, + max_frame_size: usize, + expect_masked: bool, + ) -> Result { let mut head = [0u8; 2]; - self.stream.read_exact(&mut head)?; + read.read_exact(&mut head)?; let fin = head[0] & 0x80 != 0; - let opcode = head[0] & 0x0F; - if !fin { - return Err(WebSocketError::Protocol( - "fragmented frames are not supported".into(), - )); - } + let rsv1 = head[0] & 0x40 != 0; + let rsv2 = head[0] & 0x20 != 0; + let rsv3 = head[0] & 0x10 != 0; + let opcode = Opcode::from_wire(head[0] & 0x0F)?; let masked = head[1] & 0x80 != 0; + if masked != expect_masked { + let direction = if expect_masked { "client" } else { "server" }; + return Err(WebSocketError::Protocol(format!( + "{direction} frame mask bit mismatch" + ))); + } + let mut len = (head[1] & 0x7F) as u64; if len == 126 { let mut ext = [0u8; 2]; - self.stream.read_exact(&mut ext)?; + read.read_exact(&mut ext)?; len = u16::from_be_bytes(ext) as u64; + if len < 126 { + return Err(WebSocketError::Protocol( + "non-minimal 16-bit payload length".into(), + )); + } } else if len == 127 { let mut ext = [0u8; 8]; - self.stream.read_exact(&mut ext)?; + read.read_exact(&mut ext)?; len = u64::from_be_bytes(ext); + if len < 65_536 { + return Err(WebSocketError::Protocol( + "non-minimal 64-bit payload length".into(), + )); + } + if len & (1 << 63) != 0 { + return Err(WebSocketError::Protocol( + "payload length uses forbidden high bit".into(), + )); + } } - if len > usize::MAX as u64 { + if len > max_frame_size as u64 || len > usize::MAX as u64 { return Err(WebSocketError::MessageTooLarge); } + let mask = if masked { let mut m = [0u8; 4]; - self.stream.read_exact(&mut m)?; + read.read_exact(&mut m)?; Some(m) } else { None }; let mut payload = vec![0; len as usize]; - self.stream.read_exact(&mut payload)?; + read.read_exact(&mut payload)?; if let Some(mask) = mask { - for (i, b) in payload.iter_mut().enumerate() { - *b ^= mask[i % 4]; + apply_mask(&mut payload, mask); + } + + let frame = Self { + fin, + rsv1, + rsv2, + rsv3, + opcode, + masked, + payload, + }; + frame.validate(max_frame_size, expect_masked)?; + Ok(frame) + } + + pub fn write_to(&self, write: &mut W, mask: Option<[u8; 4]>) -> Result<()> { + self.validate(usize::MAX, mask.is_some())?; + let mut out = Vec::with_capacity(self.payload.len() + 14); + let mut b0 = self.opcode.wire(); + if self.fin { + b0 |= 0x80; + } + if self.rsv1 { + b0 |= 0x40; + } + if self.rsv2 { + b0 |= 0x20; + } + if self.rsv3 { + b0 |= 0x10; + } + out.push(b0); + + let mask_bit = if mask.is_some() { 0x80 } else { 0 }; + match self.payload.len() { + 0..=125 => out.push(mask_bit | self.payload.len() as u8), + 126..=65_535 => { + out.push(mask_bit | 126); + out.extend_from_slice(&(self.payload.len() as u16).to_be_bytes()); } + _ => { + out.push(mask_bit | 127); + out.extend_from_slice(&(self.payload.len() as u64).to_be_bytes()); + } + } + + if let Some(mask) = mask { + out.extend_from_slice(&mask); + let mut payload = self.payload.clone(); + apply_mask(&mut payload, mask); + out.extend_from_slice(&payload); + } else { + out.extend_from_slice(&self.payload); } - Ok(Frame { opcode, payload }) + write.write_all(&out)?; + Ok(()) + } + + fn validate(&self, max_frame_size: usize, masked: bool) -> Result<()> { + if self.payload.len() > max_frame_size { + return Err(WebSocketError::MessageTooLarge); + } + if self.masked != masked { + return Err(WebSocketError::Protocol("frame mask bit mismatch".into())); + } + if self.rsv2 || self.rsv3 { + return Err(WebSocketError::Protocol("unsupported reserved bit".into())); + } + if self.rsv1 && !matches!(self.opcode, Opcode::Text | Opcode::Binary) { + return Err(WebSocketError::Protocol( + "RSV1 is only valid on first compressed data frame".into(), + )); + } + if self.opcode.is_control() { + if !self.fin { + return Err(WebSocketError::Protocol( + "control frames must not be fragmented".into(), + )); + } + if self.payload.len() > MAX_CONTROL_PAYLOAD { + return Err(WebSocketError::Protocol( + "control frame payload too large".into(), + )); + } + } + if self.opcode == Opcode::Continuation && self.rsv1 { + return Err(WebSocketError::Protocol( + "continuation frame must not set RSV1".into(), + )); + } + Ok(()) } } -struct Frame { - opcode: u8, - payload: Vec, +fn build_handshake_request( + url: &Url, + protocols: &[String], + config: &WebSocketConfig, + key: &str, +) -> Result { + let host = url + .host_str() + .ok_or_else(|| WebSocketError::InvalidUrl("missing host".to_string()))?; + let port = url + .port_or_default() + .ok_or_else(|| WebSocketError::InvalidUrl("missing port".to_string()))?; + let host_header = if url.port().is_some() { + format!("{host}:{port}") + } else { + host.clone() + }; + + let mut request = format!( + "GET {} HTTP/1.1\r\n\ + Host: {host_header}\r\n\ + Upgrade: websocket\r\n\ + Connection: Upgrade\r\n\ + Sec-WebSocket-Key: {key}\r\n\ + Sec-WebSocket-Version: 13\r\n", + request_target(url) + ); + if !protocols.is_empty() { + request.push_str("Sec-WebSocket-Protocol: "); + request.push_str(&protocols.join(", ")); + request.push_str("\r\n"); + } + if let Some(deflate) = &config.permessage_deflate { + request.push_str("Sec-WebSocket-Extensions: "); + request.push_str(&deflate.offer_header()); + request.push_str("\r\n"); + } + request.push_str("\r\n"); + Ok(request) +} + +fn validate_handshake_response(status: u16, headers: &[(String, String)], key: &str) -> Result<()> { + if status != 101 { + return Err(WebSocketError::Handshake(format!( + "expected status 101, got {status}" + ))); + } + if !header_token_contains(headers, "upgrade", "websocket") { + return Err(WebSocketError::Handshake("missing Upgrade header".into())); + } + if !header_token_contains(headers, "connection", "upgrade") { + return Err(WebSocketError::Handshake( + "missing Connection header".into(), + )); + } + let expected_accept = websocket_accept(key); + let accept = header_value(headers, "sec-websocket-accept").unwrap_or_default(); + if accept.trim() != expected_accept { + return Err(WebSocketError::Handshake( + "invalid Sec-WebSocket-Accept".into(), + )); + } + Ok(()) +} + +fn websocket_accept(key: &str) -> String { + base64_encode(&sha1::sha1(format!("{key}{GUID}").as_bytes())) } fn request_target(url: &Url) -> String { let mut path = url.path(); + if path.is_empty() { + path.push('/'); + } if let Some(query) = url.query() { path.push('?'); path.push_str(query); @@ -353,12 +758,46 @@ fn parse_close_payload(payload: &[u8]) -> Result<(Option, String)> { return Err(WebSocketError::Protocol("truncated close code".into())); } let code = u16::from_be_bytes([payload[0], payload[1]]); + validate_close_code(code)?; let reason = String::from_utf8(payload[2..].to_vec()) .map_err(|_| WebSocketError::Protocol("invalid close reason UTF-8".into()))?; Ok((Some(code), reason)) } -fn read_http_upgrade_response(stream: &mut WsStream) -> Result<(u16, Vec<(String, String)>)> { +fn close_payload(code: Option, reason: &str) -> Result> { + let mut payload = Vec::new(); + if let Some(code) = code { + validate_close_code(code)?; + payload.extend_from_slice(&code.to_be_bytes()); + payload.extend_from_slice(reason.as_bytes()); + } else if !reason.is_empty() { + return Err(WebSocketError::Protocol( + "close reason requires a close code".into(), + )); + } + if payload.len() > MAX_CONTROL_PAYLOAD { + return Err(WebSocketError::Protocol( + "close frame payload too large".into(), + )); + } + Ok(payload) +} + +fn validate_close_code(code: u16) -> Result<()> { + let valid = matches!( + code, + 1000 | 1001 | 1002 | 1003 | 1007 | 1008 | 1009 | 1010 | 1011 | 3000..=4999 + ); + if valid { + Ok(()) + } else { + Err(WebSocketError::Protocol(format!( + "invalid close code {code}" + ))) + } +} + +fn read_http_upgrade_response(stream: &mut R) -> Result<(u16, Vec<(String, String)>)> { let mut buf = Vec::new(); let mut byte = [0u8; 1]; while !buf.ends_with(b"\r\n\r\n") { @@ -376,7 +815,14 @@ fn read_http_upgrade_response(stream: &mut WsStream) -> Result<(u16, Vec<(String .next() .ok_or_else(|| WebSocketError::Handshake("missing status line".into()))?; let mut status_parts = status_line.split_whitespace(); - let _version = status_parts.next(); + let version = status_parts + .next() + .ok_or_else(|| WebSocketError::Handshake("malformed status line".into()))?; + if !version.starts_with("HTTP/1.") { + return Err(WebSocketError::Handshake( + "unsupported HTTP response version".into(), + )); + } let status = status_parts .next() .and_then(|s| s.parse::().ok()) @@ -387,9 +833,15 @@ fn read_http_upgrade_response(stream: &mut WsStream) -> Result<(u16, Vec<(String if line.is_empty() { continue; } - if let Some((name, value)) = line.split_once(':') { - headers.push((name.trim().to_ascii_lowercase(), value.trim().to_string())); + if line.starts_with(' ') || line.starts_with('\t') { + return Err(WebSocketError::Handshake( + "obsolete folded header line".into(), + )); } + let (name, value) = line + .split_once(':') + .ok_or_else(|| WebSocketError::Handshake("malformed header line".into()))?; + headers.push((name.trim().to_ascii_lowercase(), value.trim().to_string())); } Ok((status, headers)) } @@ -402,11 +854,86 @@ fn header_value(headers: &[(String, String)], name: &str) -> Option { } fn header_token_contains(headers: &[(String, String)], name: &str, token: &str) -> bool { - header_value(headers, name).is_some_and(|value| { - value - .split(',') - .any(|part| part.trim().eq_ignore_ascii_case(token)) - }) + headers + .iter() + .filter(|(n, _)| n.eq_ignore_ascii_case(name)) + .flat_map(|(_, v)| v.split(',')) + .any(|part| part.trim().eq_ignore_ascii_case(token)) +} + +fn negotiate_permessage_deflate( + offered: Option<&PermessageDeflateConfig>, + header: &str, +) -> Result> { + if header.trim().is_empty() { + return Ok(None); + } + let Some(offered) = offered else { + return Err(WebSocketError::Handshake( + "server selected unsolicited extension".into(), + )); + }; + + for extension in header.split(',') { + let mut parts = extension.split(';').map(str::trim); + let Some(name) = parts.next() else { + continue; + }; + if !name.eq_ignore_ascii_case("permessage-deflate") { + continue; + } + let mut negotiated = offered.clone(); + for part in parts { + if part.eq_ignore_ascii_case("client_no_context_takeover") { + negotiated.client_no_context_takeover = true; + } else if part.eq_ignore_ascii_case("server_no_context_takeover") { + negotiated.server_no_context_takeover = true; + } else if let Some(bits) = part.strip_prefix("client_max_window_bits=") { + negotiated.client_max_window_bits = Some(parse_window_bits(bits)?); + } else if let Some(bits) = part.strip_prefix("server_max_window_bits=") { + negotiated.server_max_window_bits = Some(parse_window_bits(bits)?); + } else if part.eq_ignore_ascii_case("client_max_window_bits") { + negotiated.client_max_window_bits = Some(15); + } else { + return Err(WebSocketError::Handshake(format!( + "unsupported permessage-deflate parameter {part}" + ))); + } + } + return Ok(Some(negotiated)); + } + + Err(WebSocketError::Handshake( + "unsupported Sec-WebSocket-Extensions response".into(), + )) +} + +fn parse_window_bits(bits: &str) -> Result { + let bits = bits + .trim_matches('"') + .parse::() + .map_err(|_| WebSocketError::Handshake("invalid window bits".into()))?; + if (8..=15).contains(&bits) { + Ok(bits) + } else { + Err(WebSocketError::Handshake( + "window bits must be in 8..=15".into(), + )) + } +} + +pub fn compress_permessage_deflate(payload: &[u8]) -> Vec { + deflate::deflate_fixed(payload) +} + +pub fn decompress_permessage_deflate(payload: &[u8]) -> Result> { + deflate::inflate(payload).map_err(WebSocketError::from) +} + +fn apply_mask(payload: &mut [u8], mask: [u8; 4]) { + for (i, b) in payload.iter_mut().enumerate() { + *b ^= mask[i % 4]; + } } fn random_bytes_16() -> [u8; 16] { @@ -463,130 +990,189 @@ fn base64_encode(bytes: &[u8]) -> String { out } -fn sha1(input: &[u8]) -> [u8; 20] { - let mut h0: u32 = 0x67452301; - let mut h1: u32 = 0xEFCDAB89; - let mut h2: u32 = 0x98BADCFE; - let mut h3: u32 = 0x10325476; - let mut h4: u32 = 0xC3D2E1F0; - - let bit_len = (input.len() as u64) * 8; - let mut msg = input.to_vec(); - msg.push(0x80); - while msg.len() % 64 != 56 { - msg.push(0); - } - msg.extend_from_slice(&bit_len.to_be_bytes()); - - for chunk in msg.chunks_exact(64) { - let mut w = [0u32; 80]; - for (i, word) in w.iter_mut().take(16).enumerate() { - let j = i * 4; - *word = u32::from_be_bytes([chunk[j], chunk[j + 1], chunk[j + 2], chunk[j + 3]]); - } - for i in 16..80 { - w[i] = (w[i - 3] ^ w[i - 8] ^ w[i - 14] ^ w[i - 16]).rotate_left(1); - } - let mut a = h0; - let mut b = h1; - let mut c = h2; - let mut d = h3; - let mut e = h4; - for (i, &wi) in w.iter().enumerate() { - let (f, k) = match i { - 0..=19 => ((b & c) | ((!b) & d), 0x5A827999), - 20..=39 => (b ^ c ^ d, 0x6ED9EBA1), - 40..=59 => ((b & c) | (b & d) | (c & d), 0x8F1BBCDC), - _ => (b ^ c ^ d, 0xCA62C1D6), - }; - let temp = a - .rotate_left(5) - .wrapping_add(f) - .wrapping_add(e) - .wrapping_add(k) - .wrapping_add(wi); - e = d; - d = c; - c = b.rotate_left(30); - b = a; - a = temp; - } - h0 = h0.wrapping_add(a); - h1 = h1.wrapping_add(b); - h2 = h2.wrapping_add(c); - h3 = h3.wrapping_add(d); - h4 = h4.wrapping_add(e); - } - - let mut out = [0u8; 20]; - out[..4].copy_from_slice(&h0.to_be_bytes()); - out[4..8].copy_from_slice(&h1.to_be_bytes()); - out[8..12].copy_from_slice(&h2.to_be_bytes()); - out[12..16].copy_from_slice(&h3.to_be_bytes()); - out[16..20].copy_from_slice(&h4.to_be_bytes()); - out -} - #[cfg(test)] mod tests { use super::*; - use std::io::{Read, Write}; + use std::io::{Cursor, Read, Write}; use std::net::TcpListener; use std::thread; #[test] fn sha1_and_base64_match_websocket_accept_vector() { - let accept = base64_encode(&sha1( - b"dGhlIHNhbXBsZSBub25jZQ==258EAFA5-E914-47DA-95CA-C5AB0DC85B11", + assert_eq!( + websocket_accept("dGhlIHNhbXBsZSBub25jZQ=="), + "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" + ); + } + + #[test] + fn validates_good_and_bad_handshake_accept() { + let key = "dGhlIHNhbXBsZSBub25jZQ=="; + let headers = vec![ + ("upgrade".into(), "websocket".into()), + ("connection".into(), "keep-alive, Upgrade".into()), + ("sec-websocket-accept".into(), websocket_accept(key)), + ]; + validate_handshake_response(101, &headers, key).unwrap(); + + let mut bad = headers; + bad[2].1 = "bad".into(); + assert!(matches!( + validate_handshake_response(101, &bad, key), + Err(WebSocketError::Handshake(_)) )); - assert_eq!(accept, "s3pPLMBiTxaQ9kYGzzhZRbK+xOo="); } #[test] - fn ws_echo_roundtrip() { + fn builds_handshake_with_protocols_and_extensions() { + let url = Url::parse("ws://example.test/chat?q=1").unwrap(); + let config = WebSocketConfig { + permessage_deflate: Some(PermessageDeflateConfig::default()), + ..WebSocketConfig::default() + }; + let request = build_handshake_request(&url, &["chat".into()], &config, "key").unwrap(); + assert!(request.starts_with("GET /chat?q=1 HTTP/1.1\r\n")); + assert!(request.contains("Sec-WebSocket-Protocol: chat\r\n")); + assert!(request.contains("Sec-WebSocket-Extensions: permessage-deflate")); + } + + #[test] + fn frame_serialize_parse_all_opcodes_and_lengths() { + let cases = [ + (Opcode::Continuation, 0usize), + (Opcode::Text, 5), + (Opcode::Binary, 126), + (Opcode::Close, 2), + (Opcode::Ping, 3), + (Opcode::Pong, 125), + (Opcode::Binary, 65_536), + ]; + for (opcode, len) in cases { + let payload = vec![b'x'; len]; + let frame = Frame { + fin: true, + rsv1: false, + rsv2: false, + rsv3: false, + opcode, + masked: false, + payload: payload.clone(), + }; + let mut bytes = Vec::new(); + frame.write_to(&mut bytes, None).unwrap(); + let parsed = Frame::read_from(&mut Cursor::new(bytes), len.max(125), false).unwrap(); + assert_eq!(parsed.opcode, opcode); + assert_eq!(parsed.payload, payload); + } + } + + #[test] + fn masking_roundtrip() { + let frame = Frame { + fin: true, + rsv1: false, + rsv2: false, + rsv3: false, + opcode: Opcode::Text, + masked: true, + payload: b"hello".to_vec(), + }; + let mut bytes = Vec::new(); + frame.write_to(&mut bytes, Some([1, 2, 3, 4])).unwrap(); + assert_ne!(&bytes[6..], b"hello"); + let parsed = Frame::read_from(&mut Cursor::new(bytes), 1024, true).unwrap(); + assert_eq!(parsed.payload, b"hello"); + } + + #[test] + fn rejects_invalid_control_frames() { + let frame = Frame { + fin: false, + rsv1: false, + rsv2: false, + rsv3: false, + opcode: Opcode::Ping, + masked: false, + payload: Vec::new(), + }; + assert!(matches!( + frame.validate(1024, false), + Err(WebSocketError::Protocol(_)) + )); + } + + #[test] + fn close_payload_validates_codes_and_reason() { + let payload = close_payload(Some(1000), "bye").unwrap(); + assert_eq!( + parse_close_payload(&payload).unwrap(), + (Some(1000), "bye".into()) + ); + assert!(close_payload(Some(1006), "").is_err()); + assert!(close_payload(None, "bad").is_err()); + } + + #[test] + fn permessage_deflate_roundtrip() { + let payload = b"hello hello hello hello"; + let compressed = compress_permessage_deflate(payload); + assert_ne!(compressed, payload); + assert_eq!(decompress_permessage_deflate(&compressed).unwrap(), payload); + } + + #[test] + fn permessage_deflate_negotiates_context_takeover_parameters() { + let offered = PermessageDeflateConfig::default(); + let negotiated = negotiate_permessage_deflate( + Some(&offered), + "permessage-deflate; server_no_context_takeover; client_max_window_bits=12", + ) + .unwrap() + .unwrap(); + assert!(negotiated.server_no_context_takeover); + assert_eq!(negotiated.client_max_window_bits, Some(12)); + } + + #[test] + fn ws_echo_roundtrip_fragmentation_ping_and_close() { let listener = TcpListener::bind("127.0.0.1:0").unwrap(); let addr = listener.local_addr().unwrap(); let server = thread::spawn(move || { let (mut stream, _) = listener.accept().unwrap(); - let mut req = Vec::new(); - let mut b = [0u8; 1]; - while !req.ends_with(b"\r\n\r\n") { - stream.read_exact(&mut b).unwrap(); - req.push(b[0]); - } - let req_text = String::from_utf8_lossy(&req); - let key = req_text - .lines() - .find_map(|line| { - line.split_once(':') - .filter(|(n, _)| n.eq_ignore_ascii_case("Sec-WebSocket-Key")) - }) - .map(|(_, v)| v.trim()) - .unwrap(); - let accept = base64_encode(&sha1(format!("{key}{GUID}").as_bytes())); - write!( - stream, - "HTTP/1.1 101 Switching Protocols\r\n\ - Upgrade: websocket\r\n\ - Connection: Upgrade\r\n\ - Sec-WebSocket-Accept: {accept}\r\n\r\n" - ) - .unwrap(); - let frame = read_client_frame_for_test(&mut stream); - write_server_frame_for_test(&mut stream, 0x1, &frame); - let _ = read_client_frame_for_test(&mut stream); - write_server_frame_for_test(&mut stream, 0x8, &1000u16.to_be_bytes()); + accept_handshake_for_test(&mut stream, ""); + let frame = Frame::read_from(&mut stream, 1024, true).unwrap(); + assert_eq!(frame.payload, b"hello"); + + write_server_frame_for_test(&mut stream, false, Opcode::Text, b"he"); + write_server_frame_for_test(&mut stream, true, Opcode::Ping, b"p"); + let pong = Frame::read_from(&mut stream, 1024, true).unwrap(); + assert_eq!(pong.opcode, Opcode::Pong); + assert_eq!(pong.payload, b"p"); + write_server_frame_for_test(&mut stream, true, Opcode::Continuation, b"llo"); + + let close = Frame { + fin: true, + rsv1: false, + rsv2: false, + rsv3: false, + opcode: Opcode::Close, + masked: false, + payload: 1000u16.to_be_bytes().to_vec(), + }; + close.write_to(&mut stream, None).unwrap(); + let echoed = Frame::read_from(&mut stream, 1024, true).unwrap(); + assert_eq!(echoed.opcode, Opcode::Close); }); let url = Url::parse(&format!("ws://{addr}/chat")).unwrap(); let (mut client, hs) = WebSocketClient::connect(&url, &[]).unwrap(); assert_eq!(hs.protocol, ""); client.send_text("hello").unwrap(); + assert_eq!(client.read_message().unwrap(), Message::Ping(b"p".to_vec())); assert_eq!( client.read_message().unwrap(), Message::Text("hello".into()) ); - client.close(Some(1000), "").unwrap(); assert_eq!( client.read_message().unwrap(), Message::Close { @@ -594,32 +1180,119 @@ mod tests { reason: String::new() } ); + assert!(client.close_handshake_complete()); server.join().unwrap(); } - fn read_client_frame_for_test(stream: &mut std::net::TcpStream) -> Vec { - let mut head = [0u8; 2]; - stream.read_exact(&mut head).unwrap(); - let mut len = (head[1] & 0x7f) as usize; - if len == 126 { - let mut ext = [0u8; 2]; - stream.read_exact(&mut ext).unwrap(); - len = u16::from_be_bytes(ext) as usize; - } - let mut mask = [0u8; 4]; - stream.read_exact(&mut mask).unwrap(); - let mut payload = vec![0; len]; - stream.read_exact(&mut payload).unwrap(); - for (i, b) in payload.iter_mut().enumerate() { - *b ^= mask[i % 4]; + #[test] + fn invalid_utf8_text_triggers_1007_close() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let addr = listener.local_addr().unwrap(); + let server = thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + accept_handshake_for_test(&mut stream, ""); + write_server_frame_for_test(&mut stream, true, Opcode::Text, &[0xff]); + let close = Frame::read_from(&mut stream, 1024, true).unwrap(); + assert_eq!(close.opcode, Opcode::Close); + assert_eq!(&close.payload[..2], &1007u16.to_be_bytes()); + }); + + let url = Url::parse(&format!("ws://{addr}/chat")).unwrap(); + let (mut client, _) = WebSocketClient::connect(&url, &[]).unwrap(); + assert!(matches!( + client.read_message(), + Err(WebSocketError::Protocol(_)) + )); + server.join().unwrap(); + } + + #[test] + fn compressed_message_roundtrips_over_socket() { + let listener = TcpListener::bind("127.0.0.1:0").unwrap(); + let addr = listener.local_addr().unwrap(); + let server = thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + accept_handshake_for_test( + &mut stream, + "Sec-WebSocket-Extensions: permessage-deflate; server_no_context_takeover\r\n", + ); + let frame = Frame::read_from(&mut stream, 4096, true).unwrap(); + assert!(frame.rsv1); + let payload = decompress_permessage_deflate(&frame.payload).unwrap(); + assert_eq!(payload, b"hello hello hello"); + let compressed = compress_permessage_deflate(b"hello hello hello"); + let frame = Frame { + fin: true, + rsv1: true, + rsv2: false, + rsv3: false, + opcode: Opcode::Text, + masked: false, + payload: compressed, + }; + frame.write_to(&mut stream, None).unwrap(); + }); + + let url = Url::parse(&format!("ws://{addr}/chat")).unwrap(); + let config = WebSocketConfig { + permessage_deflate: Some(PermessageDeflateConfig::default()), + ..WebSocketConfig::default() + }; + let (mut client, hs) = WebSocketClient::connect_with_config(&url, &[], config).unwrap(); + assert!(hs.extensions.contains("permessage-deflate")); + client.send_text("hello hello hello").unwrap(); + assert_eq!( + client.read_message().unwrap(), + Message::Text("hello hello hello".into()) + ); + server.join().unwrap(); + } + + fn accept_handshake_for_test(stream: &mut std::net::TcpStream, extra_headers: &str) { + let mut req = Vec::new(); + let mut b = [0u8; 1]; + while !req.ends_with(b"\r\n\r\n") { + stream.read_exact(&mut b).unwrap(); + req.push(b[0]); } - payload + let req_text = String::from_utf8_lossy(&req); + let key = req_text + .lines() + .find_map(|line| { + line.split_once(':') + .filter(|(n, _)| n.eq_ignore_ascii_case("Sec-WebSocket-Key")) + }) + .map(|(_, v)| v.trim()) + .unwrap(); + write!( + stream, + "HTTP/1.1 101 Switching Protocols\r\n\ + Upgrade: websocket\r\n\ + Connection: Upgrade\r\n\ + Sec-WebSocket-Accept: {}\r\n{}\ + \r\n", + websocket_accept(key), + extra_headers + ) + .unwrap(); } - fn write_server_frame_for_test(stream: &mut std::net::TcpStream, opcode: u8, payload: &[u8]) { - let mut frame = vec![0x80 | opcode]; - frame.push(payload.len() as u8); - frame.extend_from_slice(payload); - stream.write_all(&frame).unwrap(); + fn write_server_frame_for_test( + stream: &mut std::net::TcpStream, + fin: bool, + opcode: Opcode, + payload: &[u8], + ) { + Frame { + fin, + rsv1: false, + rsv2: false, + rsv3: false, + opcode, + masked: false, + payload: payload.to_vec(), + } + .write_to(stream, None) + .unwrap(); } } -- 2.51.2