//! WebSocket frames, adversarially (M4a spec, section 6). Everything a server sends is untrusted: //! each hostile frame must end the connection with an error, never a panic, and a length must be //! refused before anything waits for or allocates its payload. A seeded property test compares the //! decoder with a deliberately naive one written here, on valid streams and on random mutations of //! them, fed in random pieces. The seed is printed on failure. Do not edit. use gatewayd::ws::WsError; use gatewayd::ws::frame::{ CLOSE, CONTINUATION, Decoder, Incoming, MAX_MESSAGE, PING, PONG, TEXT, encode, }; /// How a server frame's length is written: the shortest form, or a longer one on purpose. #[derive(Clone, Copy)] enum Len { Short, Force16, Force64, } /// A frame as a server sends it (unmasked unless `masked`). fn frame(fin: bool, rsv: u8, opcode: u8, masked: bool, payload: &[u8], form: Len) -> Vec { let mut out = vec![(if fin { 0x80 } else { 0 }) | (rsv << 4) | opcode]; let m = if masked { 0x80 } else { 0 }; let len = payload.len(); match form { Len::Short if len < 126 => out.push(m | len as u8), Len::Short if len <= 0xFFFF => { out.push(m | 126); out.extend_from_slice(&(len as u16).to_be_bytes()); } Len::Force16 => { out.push(m | 126); out.extend_from_slice(&(len as u16).to_be_bytes()); } _ => { out.push(m | 127); out.extend_from_slice(&(len as u64).to_be_bytes()); } } if masked { out.extend_from_slice(&[1, 2, 3, 4]); } out.extend_from_slice(payload); out } fn text(s: &str) -> Vec { frame(true, 0, TEXT, false, s.as_bytes(), Len::Short) } /// Feed `bytes` in pieces of `step` and collect every message, stopping at the first error. fn decode(bytes: &[u8], step: usize) -> (Vec, Option) { let mut d = Decoder::new(); let mut got = Vec::new(); for piece in bytes.chunks(step.max(1)) { d.feed(piece); loop { match d.next_message() { Ok(Some(m)) => got.push(m), Ok(None) => break, Err(e) => return (got, Some(e.to_string())), } } } (got, None) } fn fails(bytes: &[u8], why: &str) { for step in [1, 2, 3, 7, bytes.len().max(1)] { let (_, err) = decode(bytes, step); assert!(err.is_some(), "{why} (fed {step} at a time) was accepted"); } } #[test] fn plain_messages() { let mut bytes = text("hello"); bytes.extend(frame(true, 0, PING, false, b"p1", Len::Short)); bytes.extend(frame(true, 0, PONG, false, b"", Len::Short)); bytes.extend(frame( true, 0, CLOSE, false, &[0x03, 0xE8, b'b', b'y', b'e'], Len::Short, )); let (got, err) = decode(&bytes, bytes.len()); assert_eq!(err, None); assert_eq!( got, [ Incoming::Text("hello".into()), Incoming::Ping(b"p1".to_vec()), Incoming::Pong(Vec::new()), Incoming::Close(Some(1000), "bye".into()), ] ); assert_eq!( decode(&frame(true, 0, CLOSE, false, b"", Len::Short), 1).0, [Incoming::Close(None, String::new())] ); } #[test] fn fragments_reassemble_with_control_frames_between_and_utf8_split_across_them() { let snow = "snow ☃ man"; let bytes_of = snow.as_bytes(); let cut = snow.find('☃').unwrap() + 1; // inside the three-byte character let mut bytes = frame(false, 0, TEXT, false, &bytes_of[..cut], Len::Short); bytes.extend(frame(true, 0, PING, false, b"mid", Len::Short)); bytes.extend(frame( false, 0, CONTINUATION, false, &bytes_of[cut..cut + 1], Len::Short, )); bytes.extend(frame( true, 0, CONTINUATION, false, &bytes_of[cut + 1..], Len::Short, )); for step in 1..=bytes.len() { let (got, err) = decode(&bytes, step); assert_eq!(err, None, "step {step}"); assert_eq!( got, [Incoming::Ping(b"mid".to_vec()), Incoming::Text(snow.into())], "step {step}" ); } } #[test] fn lengths_in_every_form() { for len in [0usize, 1, 125, 126, 127, 65_535, 65_536, 100_000] { let body = "x".repeat(len); let (got, err) = decode(&text(&body), 4096); assert_eq!(err, None, "{len}"); assert_eq!(got, [Incoming::Text(body)], "{len}"); } } #[test] fn hostile_frames_end_the_connection() { for rsv in [1, 2, 4] { fails( &frame(true, rsv, TEXT, false, b"x", Len::Short), "a reserved bit", ); } fails( &frame(true, 0, TEXT, true, b"x", Len::Short), "a masked frame from the server", ); for op in [2u8, 3, 7, 11, 15] { fails( &frame(true, 0, op, false, b"x", Len::Short), "an unknown or binary opcode", ); } fails( &frame(true, 0, PING, false, &[0u8; 126], Len::Short), "a control frame over 125 bytes", ); fails( &frame(false, 0, PING, false, b"x", Len::Short), "a fragmented control frame", ); fails( &frame(true, 0, CONTINUATION, false, b"x", Len::Short), "a continuation with nothing to continue", ); let mut inside = frame(false, 0, TEXT, false, b"a", Len::Short); inside.extend(text("b")); fails(&inside, "a new message inside an unfinished one"); fails( &frame(true, 0, TEXT, false, b"x", Len::Force16), "a 16-bit length for 1 byte", ); fails( &frame(true, 0, TEXT, false, &[b'y'; 200], Len::Force64), "a 64-bit length for 200 bytes", ); fails( &frame(true, 0, TEXT, false, &[0xff, 0xfe], Len::Short), "text that is not UTF-8", ); fails( &frame(true, 0, CLOSE, false, &[3], Len::Short), "a close frame of one byte", ); fails( &frame(true, 0, CLOSE, false, &[3, 232, 0xff], Len::Short), "a close reason that is not UTF-8", ); } #[test] fn huge_lengths_are_refused_from_the_header_alone() { // Only the header is fed: the decoder must refuse without waiting for a payload. let top_bit = [0x81u8, 127, 0x80, 0, 0, 0, 0, 0, 0, 1]; let mut d = Decoder::new(); d.feed(&top_bit); assert!( d.next_message().is_err(), "a 64-bit length with its top bit set" ); let too_big = (MAX_MESSAGE as u64) + 1; let mut head = vec![0x81u8, 127]; head.extend_from_slice(&too_big.to_be_bytes()); let mut d = Decoder::new(); d.feed(&head); assert!(matches!(d.next_message(), Err(WsError::TooLarge))); let mut head = vec![0x81u8, 127]; head.extend_from_slice(&0x7FFF_FFFF_FFFF_FFFFu64.to_be_bytes()); let mut d = Decoder::new(); d.feed(&head); assert!(matches!(d.next_message(), Err(WsError::TooLarge))); // Across fragments: the sum counts. let half = MAX_MESSAGE / 2 + 1; let mut d = Decoder::new(); d.feed(&frame(false, 0, TEXT, false, &vec![b'a'; half], Len::Short)); assert!(matches!(d.next_message(), Ok(None))); let mut second = vec![0x00u8, 127]; second.extend_from_slice(&(half as u64).to_be_bytes()); d.feed(&second); assert!(matches!(d.next_message(), Err(WsError::TooLarge))); // Exactly the limit is fine. let (got, err) = decode(&text(&"z".repeat(MAX_MESSAGE)), 65_536); assert_eq!(err, None); assert_eq!(got.len(), 1); } #[test] fn our_frames_are_masked_and_decode_back() { let mask = [0x11, 0x22, 0x33, 0x44]; for len in [0usize, 5, 125, 126, 65_535, 65_536] { let payload: Vec = (0..len).map(|i| (i % 251) as u8).collect(); let bytes = encode(TEXT, &payload, mask); assert_eq!(bytes[0], 0x80 | TEXT, "FIN and the opcode"); assert_ne!(bytes[1] & 0x80, 0, "the mask bit"); let (len_field, header) = match bytes[1] & 0x7F { 126 => (u16::from_be_bytes([bytes[2], bytes[3]]) as usize, 4), 127 => ( u64::from_be_bytes(bytes[2..10].try_into().unwrap()) as usize, 10, ), n => (n as usize, 2), }; assert_eq!(len_field, len); let shortest = if len < 126 { 2 } else if len <= 0xFFFF { 4 } else { 10 }; assert_eq!(header, shortest, "the shortest length form"); assert_eq!(&bytes[header..header + 4], &mask); let unmasked: Vec = bytes[header + 4..] .iter() .zip(mask.iter().cycle()) .map(|(b, m)| b ^ m) .collect(); assert_eq!(unmasked, payload); } assert_eq!(encode(PONG, b"p", mask)[0], 0x80 | PONG); } // ---------- the property test ---------- struct Rng(u64); impl Rng { fn next(&mut self) -> u64 { let mut x = self.0; x ^= x << 13; x ^= x >> 7; x ^= x << 17; self.0 = x; x } fn below(&mut self, n: usize) -> usize { (self.next() % n.max(1) as u64) as usize } } /// The naive decoder: the whole buffer at once, the rules written out plainly. fn naive(bytes: &[u8]) -> (Vec, bool) { let mut out = Vec::new(); let mut i = 0usize; let mut partial: Option> = None; while i < bytes.len() { if bytes.len() - i < 2 { return (out, false); } let (b0, b1) = (bytes[i], bytes[i + 1]); let (fin, rsv, op, masked, short) = ( b0 >> 7 == 1, (b0 >> 4) & 7, b0 & 15, b1 >> 7 == 1, (b1 & 127) as usize, ); if rsv != 0 || masked || ![0, 1, 8, 9, 10].contains(&op) { return (out, true); } let (hl, len) = if short == 126 { if bytes.len() - i < 4 { return (out, false); } let l = u16::from_be_bytes([bytes[i + 2], bytes[i + 3]]) as usize; if l < 126 { return (out, true); } (4, l) } else if short == 127 { if bytes.len() - i < 10 { return (out, false); } let l = u64::from_be_bytes(bytes[i + 2..i + 10].try_into().unwrap()); if l >> 63 == 1 || l <= 0xFFFF { return (out, true); } (10, l as usize) } else { (2, short) }; let control = op >= 8; if control && (!fin || len > 125) { return (out, true); } if !control { if (op == 1 && partial.is_some()) || (op == 0 && partial.is_none()) { return (out, true); } if partial.as_ref().map_or(0, |p| p.len()) + len > MAX_MESSAGE { return (out, true); } } if bytes.len() - i - hl < len { return (out, false); } let payload = bytes[i + hl..i + hl + len].to_vec(); i += hl + len; match op { 9 => out.push(Incoming::Ping(payload)), 10 => out.push(Incoming::Pong(payload)), 8 => { if payload.len() == 1 { return (out, true); } if payload.is_empty() { out.push(Incoming::Close(None, String::new())); } else { match String::from_utf8(payload[2..].to_vec()) { Ok(r) => out.push(Incoming::Close( Some(u16::from_be_bytes([payload[0], payload[1]])), r, )), Err(_) => return (out, true), } } } _ => { let mut m = if op == 1 { Vec::new() } else { partial.take().unwrap() }; m.extend_from_slice(&payload); if fin { match String::from_utf8(m) { Ok(t) => out.push(Incoming::Text(t)), Err(_) => return (out, true), } } else { partial = Some(m); } } } } (out, false) } /// A random valid stream: text messages split into random fragments, with control frames between. fn valid_stream(rng: &mut Rng) -> Vec { let mut bytes = Vec::new(); for _ in 0..1 + rng.below(6) { let len = [0, 1, 50, 125, 126, 300, 70_000][rng.below(7)]; let body: String = (0..len) .map(|k| { if (k + rng.below(3)).is_multiple_of(29) { 'é' } else { 'a' } }) .collect(); let raw = body.as_bytes(); let parts = 1 + rng.below(3); let mut cuts: Vec = (0..parts - 1).map(|_| rng.below(raw.len() + 1)).collect(); cuts.sort(); let mut start = 0; for (k, cut) in cuts .iter() .copied() .chain(std::iter::once(raw.len())) .enumerate() { let op = if k == 0 { TEXT } else { CONTINUATION }; bytes.extend(frame( k == parts - 1, 0, op, false, &raw[start..cut], Len::Short, )); start = cut; if rng.below(3) == 0 { bytes.extend(frame( true, 0, PING, false, &[rng.below(256) as u8; 3], Len::Short, )); } } } bytes } #[test] fn random_streams_agree_with_the_naive_decoder() { for case in 0..300u64 { let seed = 0x9E37_79B9_7F4A_7C15 ^ (case * 7919 + 1); let mut rng = Rng(seed); let mut bytes = valid_stream(&mut rng); if case % 2 == 1 { // Mutate: flip a few random bits, so most streams break somewhere different. for _ in 0..1 + rng.below(4) { let at = rng.below(bytes.len()); bytes[at] ^= 1 << rng.below(8); } } let (want, want_err) = naive(&bytes); let step = 1 + rng.below(4096); let (got, got_err) = decode(&bytes, step); assert_eq!(got, want, "seed {seed:#x}, step {step}: messages differ"); assert_eq!( got_err.is_some(), want_err, "seed {seed:#x}, step {step}: {got_err:?} vs naive error {want_err}" ); } }