1//! The WebSocket framing both ends speak (RFC 6455): whole text and binary messages, pings and
2//! closes, each message capped in size.
3
4use base64::Engine;
5use sha1::{Digest, Sha1};
6use std::io::{self, Read};
7
8pub const TEXT: u8 = 1;
9pub const BINARY: u8 = 2;
10pub const CLOSE: u8 = 8;
11pub const PING: u8 = 9;
12pub const PONG: u8 = 10;
13/// The longest HTTP head either end reads.
14pub const HEAD: usize = 8 << 10;
15
16#[derive(Debug, PartialEq, Eq)]
17pub 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`.
26pub 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.
32pub 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.
59pub 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`.
73pub 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.
81pub 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
91impl<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
180pub(crate) fn invalid(message: &'static str) -> io::Error {
181 io::Error::new(io::ErrorKind::InvalidData, message)
182}
183
184#[cfg(test)]
185mod 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}