1use std::{cmp::min, io};
2
3use bytes::{Buf, BufMut, BytesMut};
4use tracing::debug;
5
6use super::{
7 mask::apply_mask,
8 proto::{CloseCode, CloseReason, OpCode},
9 ProtocolError,
10};
11
12#[derive(Debug)]
14pub struct Parser;
15
16impl Parser {
17 fn parse_metadata(
18 src: &[u8],
19 server: bool,
20 ) -> Result<Option<(usize, bool, OpCode, usize, Option<[u8; 4]>)>, ProtocolError> {
21 let chunk_len = src.len();
22
23 let mut idx = 2;
24 if chunk_len < 2 {
25 return Ok(None);
26 }
27
28 let first = src[0];
29 let second = src[1];
30 let finished = first & 0x80 != 0;
31
32 if first & 0b0111_0000 != 0 {
34 return Err(ProtocolError::Io(io::Error::new(
36 io::ErrorKind::InvalidData,
37 "Received a frame with non-zero reserved bits",
38 )));
39 }
40
41 let masked = second & 0x80 != 0;
43 if !masked && server {
44 return Err(ProtocolError::UnmaskedFrame);
45 } else if masked && !server {
46 return Err(ProtocolError::MaskedFrame);
47 }
48
49 let opcode = OpCode::from(first & 0x0F);
51
52 if let OpCode::Bad = opcode {
53 return Err(ProtocolError::InvalidOpcode(first & 0x0F));
54 }
55
56 let len = second & 0x7F;
57 let length = if len == 126 {
58 if chunk_len < 4 {
59 return Ok(None);
60 }
61 let len = usize::from(u16::from_be_bytes(
62 TryFrom::try_from(&src[idx..idx + 2]).unwrap(),
63 ));
64 idx += 2;
65 len
66 } else if len == 127 {
67 if chunk_len < 10 {
68 return Ok(None);
69 }
70 let len = u64::from_be_bytes(TryFrom::try_from(&src[idx..idx + 8]).unwrap());
71 idx += 8;
72 len as usize
73 } else {
74 len as usize
75 };
76
77 let mask = if server {
78 if chunk_len < idx + 4 {
79 return Ok(None);
80 }
81
82 let mask = TryFrom::try_from(&src[idx..idx + 4]).unwrap();
83
84 idx += 4;
85
86 Some(mask)
87 } else {
88 None
89 };
90
91 Ok(Some((idx, finished, opcode, length, mask)))
92 }
93
94 pub fn parse(
96 src: &mut BytesMut,
97 server: bool,
98 max_size: usize,
99 ) -> Result<Option<(bool, OpCode, Option<BytesMut>)>, ProtocolError> {
100 let (idx, finished, opcode, length, mask) = match Parser::parse_metadata(src, server)? {
102 None => return Ok(None),
103 Some(res) => res,
104 };
105
106 let frame_len = match idx.checked_add(length) {
107 Some(len) => len,
108 None => return Err(ProtocolError::Overflow),
109 };
110
111 if src.len() < frame_len {
113 let min_length = min(length, max_size);
114 let required_cap = match idx.checked_add(min_length) {
115 Some(cap) => cap,
116 None => return Err(ProtocolError::Overflow),
117 };
118
119 if src.capacity() < required_cap {
120 src.reserve(required_cap - src.capacity());
121 }
122 return Ok(None);
123 }
124
125 src.advance(idx);
127
128 if length > max_size {
130 src.advance(length);
132 return Err(ProtocolError::Overflow);
133 }
134
135 if length == 0 {
137 return Ok(Some((finished, opcode, None)));
138 }
139
140 let mut data = src.split_to(length);
141
142 match opcode {
144 OpCode::Ping | OpCode::Pong if length > 125 => {
145 return Err(ProtocolError::InvalidLength(length));
146 }
147 OpCode::Close if length > 125 => {
148 debug!("Received close frame with payload length exceeding 125. Morphing to protocol close frame.");
149 return Ok(Some((true, OpCode::Close, None)));
150 }
151 _ => {}
152 }
153
154 if let Some(mask) = mask {
156 apply_mask(&mut data, mask);
157 }
158
159 Ok(Some((finished, opcode, Some(data))))
160 }
161
162 pub fn parse_close_payload(payload: &[u8]) -> Option<CloseReason> {
164 if payload.len() >= 2 {
165 let raw_code = u16::from_be_bytes(TryFrom::try_from(&payload[..2]).unwrap());
166 let code = CloseCode::from(raw_code);
167 let description = if payload.len() > 2 {
168 Some(String::from_utf8_lossy(&payload[2..]).into())
169 } else {
170 None
171 };
172 Some(CloseReason { code, description })
173 } else {
174 None
175 }
176 }
177
178 pub fn write_message<B: AsRef<[u8]>>(
180 dst: &mut BytesMut,
181 pl: B,
182 op: OpCode,
183 fin: bool,
184 mask: bool,
185 ) {
186 let payload = pl.as_ref();
187 let one = if fin {
188 0x80 | u8::from(op)
189 } else {
190 u8::from(op)
191 };
192 let payload_len = payload.len();
193 let (two, p_len) = if mask {
194 (0x80, payload_len + 4)
195 } else {
196 (0, payload_len)
197 };
198
199 if payload_len < 126 {
200 dst.reserve(p_len + 2);
201 dst.put_slice(&[one, two | payload_len as u8]);
202 } else if payload_len <= 65_535 {
203 dst.reserve(p_len + 4);
204 dst.put_slice(&[one, two | 126]);
205 dst.put_u16(payload_len as u16);
206 } else {
207 dst.reserve(p_len + 10);
208 dst.put_slice(&[one, two | 127]);
209 dst.put_u64(payload_len as u64);
210 };
211
212 if mask {
213 let mask = rand::random::<[u8; 4]>();
214 dst.put_slice(mask.as_ref());
215 dst.put_slice(payload.as_ref());
216 let pos = dst.len() - payload_len;
217 apply_mask(&mut dst[pos..], mask);
218 } else {
219 dst.put_slice(payload.as_ref());
220 }
221 }
222
223 #[inline]
225 pub fn write_close(dst: &mut BytesMut, reason: Option<CloseReason>, mask: bool) {
226 let payload = match reason {
227 None => Vec::new(),
228 Some(reason) => {
229 let mut payload = Into::<u16>::into(reason.code).to_be_bytes().to_vec();
230 if let Some(description) = reason.description {
231 payload.extend(description.as_bytes());
232 }
233 payload
234 }
235 };
236
237 Parser::write_message(dst, payload, OpCode::Close, true, mask)
238 }
239}
240
241#[cfg(test)]
242mod tests {
243 use bytes::Bytes;
244
245 use super::*;
246
247 struct F {
248 finished: bool,
249 opcode: OpCode,
250 payload: Bytes,
251 }
252
253 fn is_none(frm: &Result<Option<(bool, OpCode, Option<BytesMut>)>, ProtocolError>) -> bool {
254 matches!(*frm, Ok(None))
255 }
256
257 fn extract(frm: Result<Option<(bool, OpCode, Option<BytesMut>)>, ProtocolError>) -> F {
258 match frm {
259 Ok(Some((finished, opcode, payload))) => F {
260 finished,
261 opcode,
262 payload: payload
263 .map(|b| b.freeze())
264 .unwrap_or_else(|| Bytes::from("")),
265 },
266 _ => unreachable!("error"),
267 }
268 }
269
270 #[test]
271 fn test_parse() {
272 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
273 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
274
275 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
276 buf.extend(b"1");
277
278 let frame = extract(Parser::parse(&mut buf, false, 1024));
279 assert!(!frame.finished);
280 assert_eq!(frame.opcode, OpCode::Text);
281 assert_eq!(frame.payload.as_ref(), &b"1"[..]);
282 }
283
284 #[test]
285 fn test_parse_length0() {
286 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0000u8][..]);
287 let frame = extract(Parser::parse(&mut buf, false, 1024));
288 assert!(!frame.finished);
289 assert_eq!(frame.opcode, OpCode::Text);
290 assert!(frame.payload.is_empty());
291 }
292
293 #[test]
294 fn test_parse_length2() {
295 let mut buf = BytesMut::from(&[0b0000_0001u8, 126u8][..]);
296 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
297
298 let mut buf = BytesMut::from(&[0b0000_0001u8, 126u8][..]);
299 buf.extend(&[0u8, 4u8][..]);
300 buf.extend(b"1234");
301
302 let frame = extract(Parser::parse(&mut buf, false, 1024));
303 assert!(!frame.finished);
304 assert_eq!(frame.opcode, OpCode::Text);
305 assert_eq!(frame.payload.as_ref(), &b"1234"[..]);
306 }
307
308 #[test]
309 fn test_parse_length4() {
310 let mut buf = BytesMut::from(&[0b0000_0001u8, 127u8][..]);
311 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
312
313 let mut buf = BytesMut::from(&[0b0000_0001u8, 127u8][..]);
314 buf.extend(&[0u8, 0u8, 0u8, 0u8, 0u8, 0u8, 0u8, 4u8][..]);
315 buf.extend(b"1234");
316
317 let frame = extract(Parser::parse(&mut buf, false, 1024));
318 assert!(!frame.finished);
319 assert_eq!(frame.opcode, OpCode::Text);
320 assert_eq!(frame.payload.as_ref(), &b"1234"[..]);
321 }
322
323 #[test]
324 fn test_parse_frame_mask() {
325 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b1000_0001u8][..]);
326 buf.extend(b"0001");
327 buf.extend(b"1");
328
329 assert!(Parser::parse(&mut buf, false, 1024).is_err());
330
331 let frame = extract(Parser::parse(&mut buf, true, 1024));
332 assert!(!frame.finished);
333 assert_eq!(frame.opcode, OpCode::Text);
334 assert_eq!(frame.payload, Bytes::from(vec![1u8]));
335 }
336
337 #[test]
338 fn test_parse_frame_no_mask() {
339 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
340 buf.extend([1u8]);
341
342 assert!(Parser::parse(&mut buf, true, 1024).is_err());
343
344 let frame = extract(Parser::parse(&mut buf, false, 1024));
345 assert!(!frame.finished);
346 assert_eq!(frame.opcode, OpCode::Text);
347 assert_eq!(frame.payload, Bytes::from(vec![1u8]));
348 }
349
350 #[test]
351 fn test_parse_frame_with_rsv1_set() {
352 let mut buf = BytesMut::from(
356 &[
357 0b1100_0001u8, 0b1000_0001u8, 0, 0, 0, 0, b'a', ][..],
365 );
366
367 Parser::parse(&mut buf, true, 1024)
368 .expect_err("Should reject set RSV1 bit when no extension is negotiated");
369 }
370
371 #[test]
372 fn test_parse_frame_max_size() {
373 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0010u8][..]);
374 buf.extend([1u8, 1u8]);
375
376 assert!(Parser::parse(&mut buf, true, 1).is_err());
377
378 if let Err(ProtocolError::Overflow) = Parser::parse(&mut buf, false, 0) {
379 } else {
380 unreachable!("error");
381 }
382 }
383
384 #[test]
385 fn test_parse_frame_max_size_recoverability() {
386 let mut buf = BytesMut::new();
387 buf.extend([0b0000_0001u8, 0b0000_0010u8, 0b0000_0000u8, 0b0000_0000u8]);
389 buf.extend([0b0000_0010u8, 0b0000_0010u8, 0b1111_1111u8, 0b1111_1111u8]);
391
392 assert_eq!(buf.len(), 8);
393 assert!(matches!(
394 Parser::parse(&mut buf, false, 1),
395 Err(ProtocolError::Overflow)
396 ));
397 assert_eq!(buf.len(), 4);
398 let frame = extract(Parser::parse(&mut buf, false, 2));
399 assert!(!frame.finished);
400 assert_eq!(frame.opcode, OpCode::Binary);
401 assert_eq!(
402 frame.payload,
403 Bytes::from(vec![0b1111_1111u8, 0b1111_1111u8])
404 );
405 assert_eq!(buf.len(), 0);
406 }
407
408 #[test]
409 fn test_ping_frame() {
410 let mut buf = BytesMut::new();
411 Parser::write_message(&mut buf, Vec::from("data"), OpCode::Ping, true, false);
412
413 let mut v = vec![137u8, 4u8];
414 v.extend(b"data");
415 assert_eq!(&buf[..], &v[..]);
416 }
417
418 #[test]
419 fn test_pong_frame() {
420 let mut buf = BytesMut::new();
421 Parser::write_message(&mut buf, Vec::from("data"), OpCode::Pong, true, false);
422
423 let mut v = vec![138u8, 4u8];
424 v.extend(b"data");
425 assert_eq!(&buf[..], &v[..]);
426 }
427
428 #[test]
429 fn test_close_frame() {
430 let mut buf = BytesMut::new();
431 let reason = (CloseCode::Normal, "data");
432 Parser::write_close(&mut buf, Some(reason.into()), false);
433
434 let mut v = vec![136u8, 6u8, 3u8, 232u8];
435 v.extend(b"data");
436 assert_eq!(&buf[..], &v[..]);
437 }
438
439 #[test]
440 fn test_empty_close_frame() {
441 let mut buf = BytesMut::new();
442 Parser::write_close(&mut buf, None, false);
443 assert_eq!(&buf[..], &vec![0x88, 0x00][..]);
444 }
445
446 #[test]
447 fn test_parse_length_overflow() {
448 let buf: [u8; 14] = [
449 0x0a, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xeb, 0x0e, 0x8f,
450 ];
451 let mut buf = BytesMut::from(&buf[..]);
452 let result = Parser::parse(&mut buf, true, 65536);
453 assert!(matches!(result, Err(ProtocolError::Overflow)));
454 }
455}