| 1 | //! The WebSocket framing both ends speak (RFC 6455): whole text and binary messages, pings and |
| 2 | //! closes, each message capped in size. |
| 3 | |
| 4 | use base64::Engine; |
| 5 | use sha1::{Digest, Sha1}; |
| 6 | use std::io::{self, Read}; |
| 7 | |
| 8 | pub const TEXT: u8 = 1; |
| 9 | pub const BINARY: u8 = 2; |
| 10 | pub const CLOSE: u8 = 8; |
| 11 | pub const PING: u8 = 9; |
| 12 | pub const PONG: u8 = 10; |
| 13 | /// The longest HTTP head either end reads. |
| 14 | pub const HEAD: usize = 8 << 10; |
| 15 | |
| 16 | #[derive(Debug, PartialEq, Eq)] |
| 17 | pub enum Message { |
| 18 | Text(String), |
| 19 | Binary(Vec<u8>), |
| 20 | Ping(Vec<u8>), |
| 21 | Pong, |
| 22 | Close, |
| 23 | } |
| 24 | |
| 25 | /// The `Sec-WebSocket-Accept` answering a `Sec-WebSocket-Key` of `key`. |
| 26 | pub fn accept(key: &str) -> String { |
| 27 | let digest = Sha1::digest(format!("{key}258EAFA5-E914-47DA-95CA-C5AB0DC85B11")); |
| 28 | base64::engine::general_purpose::STANDARD.encode(digest) |
| 29 | } |
| 30 | |
| 31 | /// One unfragmented frame holding `payload`, masked with `mask` as a client's must be. |
| 32 | pub fn frame(opcode: u8, payload: &[u8], mask: Option<[u8; 4]>) -> Vec<u8> { |
| 33 | let mut frame = Vec::with_capacity(payload.len() + 14); |
| 34 | frame.push(0x80 | opcode); |
| 35 | let masked = if mask.is_some() { 0x80 } else { 0 }; |
| 36 | match payload.len() { |
| 37 | length @ 0..=125 => frame.push(masked | length as u8), |
| 38 | length @ 126..=0xffff => { |
| 39 | frame.push(masked | 126); |
| 40 | frame.extend_from_slice(&(length as u16).to_be_bytes()); |
| 41 | } |
| 42 | length => { |
| 43 | frame.push(masked | 127); |
| 44 | frame.extend_from_slice(&(length as u64).to_be_bytes()); |
| 45 | } |
| 46 | } |
| 47 | match mask { |
| 48 | Some(mask) => { |
| 49 | frame.extend_from_slice(&mask); |
| 50 | frame.extend(payload.iter().zip(mask.iter().cycle()).map(|(b, m)| b ^ m)); |
| 51 | } |
| 52 | None => frame.extend_from_slice(payload), |
| 53 | } |
| 54 | frame |
| 55 | } |
| 56 | |
| 57 | /// An HTTP head, through the blank line that ends it, read a byte at a time so that nothing |
| 58 | /// after it is consumed. |
| 59 | pub fn head(from: &mut impl Read) -> io::Result<String> { |
| 60 | let mut head = Vec::new(); |
| 61 | let mut byte = [0]; |
| 62 | while !head.ends_with(b"\r\n\r\n") { |
| 63 | if head.len() >= HEAD { |
| 64 | return Err(invalid("An HTTP head too long")); |
| 65 | } |
| 66 | from.read_exact(&mut byte)?; |
| 67 | head.push(byte[0]); |
| 68 | } |
| 69 | String::from_utf8(head).map_err(|_| invalid("An HTTP head that isn't UTF-8")) |
| 70 | } |
| 71 | |
| 72 | /// The value of header `name` in `head`. |
| 73 | pub fn header<'a>(head: &'a str, name: &str) -> Option<&'a str> { |
| 74 | head.lines().skip(1).find_map(|line| { |
| 75 | let (key, value) = line.split_once(':')?; |
| 76 | key.trim().eq_ignore_ascii_case(name).then(|| value.trim()) |
| 77 | }) |
| 78 | } |
| 79 | |
| 80 | /// Reads one direction of a connection as whole messages, joining fragments. |
| 81 | pub struct Reader<R> { |
| 82 | inner: R, |
| 83 | /// The largest message, in bytes. |
| 84 | most: usize, |
| 85 | /// Whether frames arrive masked, as a server reads them. |
| 86 | masked: bool, |
| 87 | /// A message's opcode and its fragments so far. |
| 88 | partial: Option<(u8, Vec<u8>)>, |
| 89 | } |
| 90 | |
| 91 | impl<R: Read> Reader<R> { |
| 92 | pub fn new(inner: R, most: usize, masked: bool) -> Self { |
| 93 | Self { |
| 94 | inner, |
| 95 | most, |
| 96 | masked, |
| 97 | partial: None, |
| 98 | } |
| 99 | } |
| 100 | |
| 101 | pub fn get_mut(&mut self) -> &mut R { |
| 102 | &mut self.inner |
| 103 | } |
| 104 | |
| 105 | /// The next message, or an error of kind `InvalidData` for one breaking the protocol or |
| 106 | /// the cap. |
| 107 | pub fn read(&mut self) -> io::Result<Message> { |
| 108 | loop { |
| 109 | let mut head = [0; 2]; |
| 110 | self.inner.read_exact(&mut head)?; |
| 111 | let (last, opcode) = (head[0] & 0x80 != 0, head[0] & 0x0f); |
| 112 | if head[0] & 0x70 != 0 || (head[1] & 0x80 != 0) != self.masked { |
| 113 | return Err(invalid( |
| 114 | "A WebSocket frame with reserved bits or the wrong mask", |
| 115 | )); |
| 116 | } |
| 117 | let length = match head[1] & 0x7f { |
| 118 | 126 => { |
| 119 | let mut length = [0; 2]; |
| 120 | self.inner.read_exact(&mut length)?; |
| 121 | u64::from(u16::from_be_bytes(length)) |
| 122 | } |
| 123 | 127 => { |
| 124 | let mut length = [0; 8]; |
| 125 | self.inner.read_exact(&mut length)?; |
| 126 | u64::from_be_bytes(length) |
| 127 | } |
| 128 | length => u64::from(length), |
| 129 | }; |
| 130 | let control = opcode & 0x08 != 0; |
| 131 | let room = match (&self.partial, control) { |
| 132 | (_, true) if !last => return Err(invalid("A fragmented control frame")), |
| 133 | (_, true) => 125, |
| 134 | (Some((_, so_far)), false) => self.most - so_far.len(), |
| 135 | (None, false) => self.most, |
| 136 | }; |
| 137 | if length > room as u64 { |
| 138 | return Err(invalid("A WebSocket message too large")); |
| 139 | } |
| 140 | let mut mask = [0; 4]; |
| 141 | if self.masked { |
| 142 | self.inner.read_exact(&mut mask)?; |
| 143 | } |
| 144 | let mut payload = vec![0; length as usize]; |
| 145 | self.inner.read_exact(&mut payload)?; |
| 146 | if self.masked { |
| 147 | for (byte, mask) in payload.iter_mut().zip(mask.iter().cycle()) { |
| 148 | *byte ^= mask; |
| 149 | } |
| 150 | } |
| 151 | let (opcode, payload) = match (opcode, &mut self.partial) { |
| 152 | (PING, _) => return Ok(Message::Ping(payload)), |
| 153 | (PONG, _) => return Ok(Message::Pong), |
| 154 | (CLOSE, _) => return Ok(Message::Close), |
| 155 | (TEXT | BINARY, None) if last => (opcode, payload), |
| 156 | (TEXT | BINARY, None) => { |
| 157 | self.partial = Some((opcode, payload)); |
| 158 | continue; |
| 159 | } |
| 160 | (0, Some((_, so_far))) => { |
| 161 | so_far.extend_from_slice(&payload); |
| 162 | if !last { |
| 163 | continue; |
| 164 | } |
| 165 | self.partial.take().expect("a partial message") |
| 166 | } |
| 167 | _ => return Err(invalid("An unexpected WebSocket frame")), |
| 168 | }; |
| 169 | return Ok(if opcode == TEXT { |
| 170 | Message::Text( |
| 171 | String::from_utf8(payload).map_err(|_| invalid("Text that isn't UTF-8"))?, |
| 172 | ) |
| 173 | } else { |
| 174 | Message::Binary(payload) |
| 175 | }); |
| 176 | } |
| 177 | } |
| 178 | } |
| 179 | |
| 180 | pub(crate) fn invalid(message: &'static str) -> io::Error { |
| 181 | io::Error::new(io::ErrorKind::InvalidData, message) |
| 182 | } |
| 183 | |
| 184 | #[cfg(test)] |
| 185 | mod tests { |
| 186 | use super::*; |
| 187 | |
| 188 | #[test] |
| 189 | fn frames_read_back_masked_or_not_and_fragments_join() { |
| 190 | let mut bytes = frame(TEXT, b"welcome 1", Some([1, 2, 3, 4])); |
| 191 | bytes.extend(frame(BINARY, &[7; 70_000], Some([9, 8, 7, 6]))); |
| 192 | bytes.extend(frame(PING, b"hi", Some([0; 4]))); |
| 193 | let mut reader = Reader::new(&bytes[..], 1 << 20, true); |
| 194 | assert_eq!(reader.read().unwrap(), Message::Text("welcome 1".into())); |
| 195 | assert_eq!(reader.read().unwrap(), Message::Binary(vec![7; 70_000])); |
| 196 | assert_eq!(reader.read().unwrap(), Message::Ping(b"hi".to_vec())); |
| 197 | |
| 198 | // "ab" then "cd" as a fragmented binary message, a ping between them. |
| 199 | let mut fragments = vec![BINARY, 2, b'a', b'b']; |
| 200 | fragments.extend(frame(PING, b"", None)); |
| 201 | fragments.extend([0x80, 2, b'c', b'd']); |
| 202 | let mut reader = Reader::new(&fragments[..], 4, false); |
| 203 | assert_eq!(reader.read().unwrap(), Message::Ping(vec![])); |
| 204 | assert_eq!(reader.read().unwrap(), Message::Binary(b"abcd".to_vec())); |
| 205 | } |
| 206 | |
| 207 | #[test] |
| 208 | fn a_message_over_the_cap_or_masked_wrongly_is_refused() { |
| 209 | let big = frame(BINARY, &[0; 10], None); |
| 210 | let error = Reader::new(&big[..], 9, false).read().unwrap_err(); |
| 211 | assert_eq!(error.kind(), io::ErrorKind::InvalidData); |
| 212 | let unmasked = frame(BINARY, b"x", None); |
| 213 | assert!(Reader::new(&unmasked[..], 9, true).read().is_err()); |
| 214 | } |
| 215 | |
| 216 | /// RFC 6455's own example. |
| 217 | #[test] |
| 218 | fn accepts_the_rfc_key() { |
| 219 | assert_eq!( |
| 220 | accept("dGhlIHNhbXBsZSBub25jZQ=="), |
| 221 | "s3pPLMBiTxaQ9kYGzzhZRbK+xOo=" |
| 222 | ); |
| 223 | } |
| 224 | } |