1use std::{cmp::min, io, str};
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 #[deprecated(
168 since = "3.13.5",
169 note = "Use `Parser::try_parse_close_payload` instead."
170 )]
171 pub fn parse_close_payload(payload: &[u8]) -> Option<CloseReason> {
172 if payload.len() >= 2 {
173 let raw_code = u16::from_be_bytes(TryFrom::try_from(&payload[..2]).unwrap());
174 let code = CloseCode::from(raw_code);
175 let description = if payload.len() > 2 {
176 Some(String::from_utf8_lossy(&payload[2..]).into())
177 } else {
178 None
179 };
180 Some(CloseReason { code, description })
181 } else {
182 None
183 }
184 }
185
186 pub fn try_parse_close_payload(payload: &[u8]) -> Result<Option<CloseReason>, ProtocolError> {
198 if payload.len() == 1 {
200 return Err(ProtocolError::InvalidLength(payload.len()));
201 }
202
203 if payload.len() >= 2 {
204 let raw_code = u16::from_be_bytes(
205 payload[..2]
206 .try_into()
207 .expect("Payload length should be checked before parsing"),
208 );
209
210 if !matches!(raw_code, 1000..=1003 | 1007..=1014 | 3000..=4999) {
214 return Err(ProtocolError::BadOpCode);
217 }
218
219 if payload.len() > 2 {
220 str::from_utf8(&payload[2..]).map_err(|_| {
223 ProtocolError::BadOpCode
226 })?;
227 }
228 }
229
230 #[expect(deprecated)]
231 Ok(Self::parse_close_payload(payload))
232 }
233
234 pub fn write_message<B: AsRef<[u8]>>(
236 dst: &mut BytesMut,
237 pl: B,
238 op: OpCode,
239 fin: bool,
240 mask: bool,
241 ) {
242 let payload = pl.as_ref();
243 let one = if fin {
244 0x80 | u8::from(op)
245 } else {
246 u8::from(op)
247 };
248 let payload_len = payload.len();
249 let (two, p_len) = if mask {
250 (0x80, payload_len + 4)
251 } else {
252 (0, payload_len)
253 };
254
255 if payload_len < 126 {
256 dst.reserve(p_len + 2);
257 dst.put_slice(&[one, two | payload_len as u8]);
258 } else if payload_len <= 65_535 {
259 dst.reserve(p_len + 4);
260 dst.put_slice(&[one, two | 126]);
261 dst.put_u16(payload_len as u16);
262 } else {
263 dst.reserve(p_len + 10);
264 dst.put_slice(&[one, two | 127]);
265 dst.put_u64(payload_len as u64);
266 };
267
268 if mask {
269 let mask = rand::random::<[u8; 4]>();
270 dst.put_slice(mask.as_ref());
271 dst.put_slice(payload.as_ref());
272 let pos = dst.len() - payload_len;
273 apply_mask(&mut dst[pos..], mask);
274 } else {
275 dst.put_slice(payload.as_ref());
276 }
277 }
278
279 #[inline]
281 pub fn write_close(dst: &mut BytesMut, reason: Option<CloseReason>, mask: bool) {
282 let payload = match reason {
283 None => Vec::new(),
284 Some(reason) => {
285 let mut payload = Into::<u16>::into(reason.code).to_be_bytes().to_vec();
286 if let Some(description) = reason.description {
287 payload.extend(description.as_bytes());
288 }
289 payload
290 }
291 };
292
293 Parser::write_message(dst, payload, OpCode::Close, true, mask)
294 }
295}
296
297#[cfg(test)]
298mod tests {
299 use bytes::Bytes;
300
301 use super::*;
302
303 struct F {
304 finished: bool,
305 opcode: OpCode,
306 payload: Bytes,
307 }
308
309 fn is_none(frm: &Result<Option<(bool, OpCode, Option<BytesMut>)>, ProtocolError>) -> bool {
310 matches!(*frm, Ok(None))
311 }
312
313 fn extract(frm: Result<Option<(bool, OpCode, Option<BytesMut>)>, ProtocolError>) -> F {
314 match frm {
315 Ok(Some((finished, opcode, payload))) => F {
316 finished,
317 opcode,
318 payload: payload
319 .map(|b| b.freeze())
320 .unwrap_or_else(|| Bytes::from("")),
321 },
322 _ => unreachable!("error"),
323 }
324 }
325
326 #[test]
327 fn test_parse() {
328 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
329 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
330
331 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
332 buf.extend(b"1");
333
334 let frame = extract(Parser::parse(&mut buf, false, 1024));
335 assert!(!frame.finished);
336 assert_eq!(frame.opcode, OpCode::Text);
337 assert_eq!(frame.payload.as_ref(), &b"1"[..]);
338 }
339
340 #[test]
341 fn test_parse_length0() {
342 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0000u8][..]);
343 let frame = extract(Parser::parse(&mut buf, false, 1024));
344 assert!(!frame.finished);
345 assert_eq!(frame.opcode, OpCode::Text);
346 assert!(frame.payload.is_empty());
347 }
348
349 #[test]
350 fn test_parse_length2() {
351 let mut buf = BytesMut::from(&[0b0000_0001u8, 126u8][..]);
352 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
353
354 let mut buf = BytesMut::from(&[0b0000_0001u8, 126u8][..]);
355 buf.extend(&[0u8, 4u8][..]);
356 buf.extend(b"1234");
357
358 let frame = extract(Parser::parse(&mut buf, false, 1024));
359 assert!(!frame.finished);
360 assert_eq!(frame.opcode, OpCode::Text);
361 assert_eq!(frame.payload.as_ref(), &b"1234"[..]);
362 }
363
364 #[test]
365 fn test_parse_length4() {
366 let mut buf = BytesMut::from(&[0b0000_0001u8, 127u8][..]);
367 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
368
369 let mut buf = BytesMut::from(&[0b0000_0001u8, 127u8][..]);
370 buf.extend(&[0u8, 0u8, 0u8, 0u8, 0u8, 0u8, 0u8, 4u8][..]);
371 buf.extend(b"1234");
372
373 let frame = extract(Parser::parse(&mut buf, false, 1024));
374 assert!(!frame.finished);
375 assert_eq!(frame.opcode, OpCode::Text);
376 assert_eq!(frame.payload.as_ref(), &b"1234"[..]);
377 }
378
379 #[test]
380 fn test_parse_frame_mask() {
381 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b1000_0001u8][..]);
382 buf.extend(b"0001");
383 buf.extend(b"1");
384
385 assert!(Parser::parse(&mut buf, false, 1024).is_err());
386
387 let frame = extract(Parser::parse(&mut buf, true, 1024));
388 assert!(!frame.finished);
389 assert_eq!(frame.opcode, OpCode::Text);
390 assert_eq!(frame.payload, Bytes::from(vec![1u8]));
391 }
392
393 #[test]
394 fn test_parse_frame_no_mask() {
395 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
396 buf.extend([1u8]);
397
398 assert!(Parser::parse(&mut buf, true, 1024).is_err());
399
400 let frame = extract(Parser::parse(&mut buf, false, 1024));
401 assert!(!frame.finished);
402 assert_eq!(frame.opcode, OpCode::Text);
403 assert_eq!(frame.payload, Bytes::from(vec![1u8]));
404 }
405
406 #[test]
407 fn test_parse_frame_with_rsv1_set() {
408 let mut buf = BytesMut::from(
412 &[
413 0b1100_0001u8, 0b1000_0001u8, 0, 0, 0, 0, b'a', ][..],
421 );
422
423 Parser::parse(&mut buf, true, 1024)
424 .expect_err("Should reject set RSV1 bit when no extension is negotiated");
425 }
426
427 #[test]
428 fn test_parse_frame_max_size() {
429 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0010u8][..]);
430 buf.extend([1u8, 1u8]);
431
432 assert!(Parser::parse(&mut buf, true, 1).is_err());
433
434 if let Err(ProtocolError::Overflow) = Parser::parse(&mut buf, false, 0) {
435 } else {
436 unreachable!("error");
437 }
438 }
439
440 #[test]
441 fn test_parse_frame_max_size_recoverability() {
442 let mut buf = BytesMut::new();
443 buf.extend([0b0000_0001u8, 0b0000_0010u8, 0b0000_0000u8, 0b0000_0000u8]);
445 buf.extend([0b0000_0010u8, 0b0000_0010u8, 0b1111_1111u8, 0b1111_1111u8]);
447
448 assert_eq!(buf.len(), 8);
449 assert!(matches!(
450 Parser::parse(&mut buf, false, 1),
451 Err(ProtocolError::Overflow)
452 ));
453 assert_eq!(buf.len(), 4);
454 let frame = extract(Parser::parse(&mut buf, false, 2));
455 assert!(!frame.finished);
456 assert_eq!(frame.opcode, OpCode::Binary);
457 assert_eq!(
458 frame.payload,
459 Bytes::from(vec![0b1111_1111u8, 0b1111_1111u8])
460 );
461 assert_eq!(buf.len(), 0);
462 }
463
464 #[test]
465 fn test_ping_frame() {
466 let mut buf = BytesMut::new();
467 Parser::write_message(&mut buf, Vec::from("data"), OpCode::Ping, true, false);
468
469 let mut v = vec![137u8, 4u8];
470 v.extend(b"data");
471 assert_eq!(&buf[..], &v[..]);
472 }
473
474 #[test]
475 fn test_pong_frame() {
476 let mut buf = BytesMut::new();
477 Parser::write_message(&mut buf, Vec::from("data"), OpCode::Pong, true, false);
478
479 let mut v = vec![138u8, 4u8];
480 v.extend(b"data");
481 assert_eq!(&buf[..], &v[..]);
482 }
483
484 #[test]
485 fn test_close_frame() {
486 let mut buf = BytesMut::new();
487 let reason = (CloseCode::Normal, "data");
488 Parser::write_close(&mut buf, Some(reason.into()), false);
489
490 let mut v = vec![136u8, 6u8, 3u8, 232u8];
491 v.extend(b"data");
492 assert_eq!(&buf[..], &v[..]);
493 }
494
495 #[test]
496 fn test_empty_close_frame() {
497 let mut buf = BytesMut::new();
498 Parser::write_close(&mut buf, None, false);
499 assert_eq!(&buf[..], &vec![0x88, 0x00][..]);
500 }
501
502 #[test]
503 fn try_parse_close_payload_validates_payload() {
504 assert!(matches!(
505 Parser::try_parse_close_payload(&[0x03, 0xe8, 0xff]).unwrap_err(),
506 ProtocolError::BadOpCode
507 ));
508 assert!(matches!(
509 Parser::try_parse_close_payload(&[0, 0]).unwrap_err(),
510 ProtocolError::BadOpCode
511 ));
512 }
513
514 #[test]
515 fn test_parse_length_overflow() {
516 let buf: [u8; 14] = [
517 0x0a, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xeb, 0x0e, 0x8f,
518 ];
519 let mut buf = BytesMut::from(&buf[..]);
520 let result = Parser::parse(&mut buf, true, 65536);
521 assert!(matches!(result, Err(ProtocolError::Overflow)));
522 }
523}