1use nanorand::Rng;
2
3use super::proto::{CloseCode, CloseReason, OpCode};
4use super::{error::ProtocolError, mask::apply_mask};
5use crate::util::{BufMut, BytePage, BytePages, Bytes, BytesMut};
6
7#[derive(Debug)]
9pub struct Parser;
10
11impl Parser {
12 fn parse_metadata(
13 src: &[u8],
14 server: bool,
15 max_size: usize,
16 ) -> Result<Option<(usize, bool, OpCode, usize, Option<u32>)>, ProtocolError> {
17 let chunk_len = src.len();
18
19 let mut idx = 2;
20 if chunk_len < 2 {
21 return Ok(None);
22 }
23
24 let first = src[0];
25 let second = src[1];
26 let finished = first & 0x80 != 0;
27 let reserved = first & 0x70;
28 if reserved != 0 {
29 return Err(ProtocolError::ReservedBits(reserved >> 4));
30 }
31
32 let masked = second & 0x80 != 0;
34 if !masked && server {
35 return Err(ProtocolError::UnmaskedFrame);
36 } else if masked && !server {
37 return Err(ProtocolError::MaskedFrame);
38 }
39
40 let raw_opcode = first & 0x0F;
42 let opcode =
43 OpCode::try_from(raw_opcode).map_err(|()| ProtocolError::InvalidOpcode(raw_opcode))?;
44 if opcode.is_control() && !finished {
45 return Err(ProtocolError::FragmentedControlFrame(opcode));
46 }
47
48 let len = second & 0x7F;
49 let length = if len == 126 {
50 if chunk_len < 4 {
51 return Ok(None);
52 }
53 let len = usize::from(u16::from_be_bytes(
54 TryFrom::try_from(&src[idx..idx + 2]).unwrap(),
55 ));
56 if len < 126 {
57 return Err(ProtocolError::InvalidLengthEncoding);
58 }
59 idx += 2;
60 len
61 } else if len == 127 {
62 if chunk_len < 10 {
63 return Ok(None);
64 }
65 if src[idx] & 0x80 != 0 {
66 return Err(ProtocolError::InvalidLengthEncoding);
67 }
68 let len = u64::from_be_bytes(TryFrom::try_from(&src[idx..idx + 8]).unwrap());
69 if len < 65_536 {
70 return Err(ProtocolError::InvalidLengthEncoding);
71 }
72 if len > max_size as u64 {
73 return Err(ProtocolError::Overflow);
74 }
75 idx += 8;
76 len as usize
77 } else {
78 len as usize
79 };
80
81 if opcode.is_control() && length > 125 {
82 return Err(ProtocolError::InvalidLength(length));
83 }
84
85 if length > max_size {
87 return Err(ProtocolError::Overflow);
88 }
89
90 let mask = if server {
91 if chunk_len < idx + 4 {
92 return Ok(None);
93 }
94
95 let mask = u32::from_ne_bytes(TryFrom::try_from(&src[idx..idx + 4]).unwrap());
96 idx += 4;
97 Some(mask)
98 } else {
99 None
100 };
101
102 Ok(Some((idx, finished, opcode, length, mask)))
103 }
104
105 pub fn parse(
115 src: &mut BytesMut,
116 server: bool,
117 max_size: usize,
118 ) -> Result<Option<(bool, OpCode, Option<Bytes>)>, ProtocolError> {
119 let Some((idx, finished, opcode, length, mask)) =
121 Parser::parse_metadata(src, server, max_size)?
122 else {
123 return Ok(None);
124 };
125
126 if src.len() < idx + length {
128 return Ok(None);
129 }
130
131 src.advance_to(idx);
133
134 if length == 0 {
136 return Ok(Some((finished, opcode, None)));
137 }
138
139 if let Some(mask) = mask {
141 apply_mask(&mut src[..length], mask);
142 }
143
144 Ok(Some((finished, opcode, Some(src.split_to(length)))))
145 }
146
147 pub fn parse_close_payload(payload: &[u8]) -> Result<Option<CloseReason>, ProtocolError> {
156 if payload.is_empty() {
157 return Ok(None);
158 }
159
160 if payload.len() == 1 {
161 return Err(ProtocolError::InvalidClosePayload);
162 }
163
164 let raw_code = u16::from_be_bytes(TryFrom::try_from(&payload[..2]).unwrap());
165 let code = CloseCode::from(raw_code);
166 if !code.is_valid() {
167 return Err(ProtocolError::InvalidCloseCode(raw_code));
168 }
169 let description = if payload.len() > 2 {
170 Some(
171 std::str::from_utf8(&payload[2..])
172 .map_err(|_| ProtocolError::InvalidUtf8)?
173 .to_owned(),
174 )
175 } else {
176 None
177 };
178 Ok(Some(CloseReason { code, description }))
179 }
180
181 pub fn write_message<B>(
191 dst: &mut BytePages,
192 pl: B,
193 op: OpCode,
194 fin: bool,
195 mask: bool,
196 ) -> Result<(), ProtocolError>
197 where
198 BytePage: From<B>,
199 {
200 let payload = BytePage::from(pl);
201 if op.is_control() {
202 if !fin {
203 return Err(ProtocolError::FragmentedControlFrame(op));
204 }
205 if payload.len() > 125 {
206 return Err(ProtocolError::InvalidLength(payload.len()));
207 }
208 }
209
210 let one: u8 = if fin {
211 0x80 | Into::<u8>::into(op)
212 } else {
213 op.into()
214 };
215 let payload_len = payload.len();
216 let two = if mask { 0x80 } else { 0 };
217
218 if payload_len < 126 {
219 dst.extend_from_slice(&[one, two | payload_len as u8]);
220 } else if payload_len <= 65_535 {
221 dst.extend_from_slice(&[one, two | 0x007e]);
222 dst.put_u16(payload_len as u16);
223 } else {
224 dst.extend_from_slice(&[one, two | 127]);
225 dst.put_u64(payload_len as u64);
226 }
227
228 if mask {
229 let mask: u32 = nanorand::tls_rng().generate();
230 let mut buf = BytesMut::from(payload);
231 apply_mask(&mut buf, mask);
232 dst.extend_from_slice(&mask.to_ne_bytes());
233 dst.append::<BytesMut>(buf);
234 } else {
235 dst.append::<BytePage>(payload);
236 }
237 Ok(())
238 }
239
240 #[inline]
247 pub fn write_close(
248 dst: &mut BytePages,
249 reason: Option<CloseReason>,
250 mask: bool,
251 ) -> Result<(), ProtocolError> {
252 let payload = match reason {
253 None => Bytes::new(),
254 Some(reason) => {
255 if !reason.code.is_valid() || (!mask && matches!(reason.code, CloseCode::Extension))
256 {
257 return Err(ProtocolError::InvalidCloseCode(reason.code.into()));
258 }
259 let mut payload =
260 BytesMut::with_capacity(reason.description.as_ref().map_or(0, String::len) + 2);
261 payload.put_u16(u16::from(reason.code));
262 if let Some(description) = reason.description {
263 payload.extend_from_slice(description.as_bytes());
264 }
265 payload.freeze()
266 }
267 };
268
269 Parser::write_message(dst, payload, OpCode::Close, true, mask)
270 }
271}
272
273#[cfg(test)]
274mod tests {
275 use super::*;
276
277 struct F {
278 finished: bool,
279 opcode: OpCode,
280 payload: Bytes,
281 }
282
283 fn is_none(frm: &Result<Option<(bool, OpCode, Option<Bytes>)>, ProtocolError>) -> bool {
284 matches!(*frm, Ok(None))
285 }
286
287 fn extract(frm: Result<Option<(bool, OpCode, Option<Bytes>)>, ProtocolError>) -> F {
288 match frm {
289 Ok(Some((finished, opcode, payload))) => F {
290 finished,
291 opcode,
292 payload: payload.unwrap_or_else(Bytes::new),
293 },
294 _ => unreachable!("error"),
295 }
296 }
297
298 #[test]
299 fn test_parse() {
300 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
301 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
302
303 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
304 buf.extend(b"1");
305
306 let frame = extract(Parser::parse(&mut buf, false, 1024));
307 assert!(!frame.finished);
308 assert_eq!(frame.opcode, OpCode::Text);
309 assert_eq!(frame.payload.as_ref(), &b"1"[..]);
310 }
311
312 #[test]
313 fn test_parse_length0() {
314 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0000u8][..]);
315 let frame = extract(Parser::parse(&mut buf, false, 1024));
316 assert!(!frame.finished);
317 assert_eq!(frame.opcode, OpCode::Text);
318 assert!(frame.payload.is_empty());
319 }
320
321 #[test]
322 fn test_reserved_bits() {
323 for reserved in [0x10, 0x20, 0x40, 0x70] {
324 let mut buf = BytesMut::from(&[0x81 | reserved, 0][..]);
325 assert!(matches!(
326 Parser::parse(&mut buf, false, 1024),
327 Err(ProtocolError::ReservedBits(_))
328 ));
329 }
330 }
331
332 #[test]
333 fn test_invalid_control_frames() {
334 let mut fragmented = BytesMut::from(&[0x09, 0][..]);
335 assert!(matches!(
336 Parser::parse(&mut fragmented, false, 1024),
337 Err(ProtocolError::FragmentedControlFrame(OpCode::Ping))
338 ));
339
340 let mut oversized = BytesMut::from(&[0x88, 126, 0, 126][..]);
341 assert!(matches!(
342 Parser::parse(&mut oversized, false, 1024),
343 Err(ProtocolError::InvalidLength(126))
344 ));
345 }
346
347 #[test]
348 fn test_parse_length2() {
349 let mut buf = BytesMut::from(&[0b0000_0001u8, 126u8][..]);
350 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
351
352 let mut buf = BytesMut::from(&[0b0000_0001u8, 126u8][..]);
353 buf.extend(&[0u8, 126u8][..]);
354 buf.extend(vec![1; 126]);
355
356 let frame = extract(Parser::parse(&mut buf, false, 1024));
357 assert!(!frame.finished);
358 assert_eq!(frame.opcode, OpCode::Text);
359 assert_eq!(frame.payload.len(), 126);
360 }
361
362 #[test]
363 fn test_parse_length4() {
364 let mut buf = BytesMut::from(&[0b0000_0001u8, 127u8][..]);
365 assert!(is_none(&Parser::parse(&mut buf, false, 1024)));
366
367 let mut buf = BytesMut::from(&[0b0000_0001u8, 127u8][..]);
368 buf.extend(&[0u8, 0u8, 0u8, 0u8, 0u8, 1u8, 0u8, 0u8][..]);
369 buf.extend(vec![1; 65_536]);
370
371 let frame = extract(Parser::parse(&mut buf, false, 65_536));
372 assert!(!frame.finished);
373 assert_eq!(frame.opcode, OpCode::Text);
374 assert_eq!(frame.payload.len(), 65_536);
375 }
376
377 #[test]
378 fn test_noncanonical_lengths() {
379 let mut short = BytesMut::from(&[0x82, 126, 0, 125][..]);
380 assert!(matches!(
381 Parser::parse(&mut short, false, usize::MAX),
382 Err(ProtocolError::InvalidLengthEncoding)
383 ));
384
385 let mut medium = BytesMut::from(&[0x82, 127, 0, 0, 0, 0, 0, 0, 0xff, 0xff][..]);
386 assert!(matches!(
387 Parser::parse(&mut medium, false, usize::MAX),
388 Err(ProtocolError::InvalidLengthEncoding)
389 ));
390
391 let mut high_bit = BytesMut::from(&[0x82, 127, 0x80, 0, 0, 0, 0, 1, 0, 0][..]);
392 assert!(matches!(
393 Parser::parse(&mut high_bit, false, usize::MAX),
394 Err(ProtocolError::InvalidLengthEncoding)
395 ));
396 }
397
398 #[test]
399 fn test_parse_frame_mask() {
400 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b1000_0001u8][..]);
401 buf.extend(b"0001");
402 buf.extend(b"1");
403
404 assert!(Parser::parse(&mut buf, false, 1024).is_err());
405
406 let frame = extract(Parser::parse(&mut buf, true, 1024));
407 assert!(!frame.finished);
408 assert_eq!(frame.opcode, OpCode::Text);
409 assert_eq!(frame.payload, Bytes::from(vec![1u8]));
410 }
411
412 #[test]
413 fn test_parse_frame_no_mask() {
414 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0001u8][..]);
415 buf.extend([1u8]);
416
417 assert!(Parser::parse(&mut buf, true, 1024).is_err());
418
419 let frame = extract(Parser::parse(&mut buf, false, 1024));
420 assert!(!frame.finished);
421 assert_eq!(frame.opcode, OpCode::Text);
422 assert_eq!(frame.payload, Bytes::from(vec![1u8]));
423 }
424
425 #[test]
426 fn test_parse_frame_max_size() {
427 let mut buf = BytesMut::from(&[0b0000_0001u8, 0b0000_0010u8][..]);
428 buf.extend([1u8, 1u8]);
429
430 assert!(Parser::parse(&mut buf, true, 1).is_err());
431
432 if let Err(ProtocolError::Overflow) = Parser::parse(&mut buf, false, 0) {
433 } else {
434 unreachable!("error");
435 }
436 }
437
438 #[test]
439 fn test_masked_frames_roundtrip_with_distinct_masks() {
440 let mut masks = Vec::new();
441 for _ in 0..4 {
442 let mut pages = BytePages::default();
443 Parser::write_message(&mut pages, Bytes::from("data"), OpCode::Binary, true, true)
444 .unwrap();
445 let mut buf = BytesMut::from(&Bytes::from(pages)[..]);
446 masks.push(buf[2..6].to_vec());
447
448 let frame = extract(Parser::parse(&mut buf, true, 1024));
449 assert!(frame.finished);
450 assert_eq!(frame.opcode, OpCode::Binary);
451 assert_eq!(frame.payload, Bytes::from("data"));
452 }
453 masks.dedup();
454 assert!(masks.len() > 1);
455 }
456
457 #[test]
458 fn test_ping_frame() {
459 let mut buf = BytePages::default();
460 Parser::write_message(&mut buf, Bytes::from("data"), OpCode::Ping, true, false).unwrap();
461
462 let mut v = vec![137u8, 4u8];
463 v.extend(b"data");
464 assert_eq!(&Bytes::from(buf)[..], &v[..]);
465 }
466
467 #[test]
468 fn test_pong_frame() {
469 let mut buf = BytePages::default();
470 Parser::write_message(&mut buf, Bytes::from("data"), OpCode::Pong, true, false).unwrap();
471
472 let mut v = vec![138u8, 4u8];
473 v.extend(b"data");
474 assert_eq!(&Bytes::from(buf)[..], &v[..]);
475 }
476
477 #[test]
478 fn test_close_frame() {
479 let mut buf = BytePages::default();
480 let reason = (CloseCode::Normal, "data");
481 Parser::write_close(&mut buf, Some(reason.into()), false).unwrap();
482
483 let mut v = vec![136u8, 6u8, 3u8, 232u8];
484 v.extend(b"data");
485 assert_eq!(&Bytes::from(buf)[..], &v[..]);
486 }
487
488 #[test]
489 fn test_empty_close_frame() {
490 let mut buf = BytePages::default();
491 Parser::write_close(&mut buf, None, false).unwrap();
492 assert_eq!(&Bytes::from(buf)[..], &[0x88, 0x00]);
493 }
494
495 #[test]
496 fn test_close_validation() {
497 assert!(matches!(
498 Parser::parse_close_payload(&[1]),
499 Err(ProtocolError::InvalidClosePayload)
500 ));
501 assert!(matches!(
502 Parser::parse_close_payload(&1006u16.to_be_bytes()),
503 Err(ProtocolError::InvalidCloseCode(1006))
504 ));
505 assert!(matches!(
506 Parser::parse_close_payload(&[0x03, 0xe8, 0xff]),
507 Err(ProtocolError::InvalidUtf8)
508 ));
509
510 let mut buf = BytePages::default();
511 assert!(matches!(
512 Parser::write_message(
513 &mut buf,
514 Bytes::from(vec![0; 126]),
515 OpCode::Ping,
516 true,
517 false
518 ),
519 Err(ProtocolError::InvalidLength(126))
520 ));
521 assert!(matches!(
522 Parser::write_message(&mut buf, Bytes::new(), OpCode::Pong, false, false),
523 Err(ProtocolError::FragmentedControlFrame(OpCode::Pong))
524 ));
525 assert!(matches!(
526 Parser::write_close(
527 &mut buf,
528 Some(CloseReason {
529 code: CloseCode::Other(2000),
530 description: None,
531 }),
532 false
533 ),
534 Err(ProtocolError::InvalidCloseCode(2000))
535 ));
536 assert!(matches!(
537 Parser::write_close(
538 &mut buf,
539 Some(CloseReason {
540 code: CloseCode::Extension,
541 description: None,
542 }),
543 false
544 ),
545 Err(ProtocolError::InvalidCloseCode(1010))
546 ));
547 assert!(matches!(
548 Parser::write_close(
549 &mut buf,
550 Some(CloseReason {
551 code: CloseCode::Normal,
552 description: Some("x".repeat(124)),
553 }),
554 false
555 ),
556 Err(ProtocolError::InvalidLength(126))
557 ));
558 }
559
560 #[test]
561 fn test_parse_edge_cases() {
562 let mut buf = BytesMut::from(&[0x82u8, 127, 0, 0, 0, 0, 0, 1, 0, 0][..]);
564 assert!(matches!(
565 Parser::parse(&mut buf, false, 65_535),
566 Err(ProtocolError::Overflow)
567 ));
568
569 let mut buf = BytesMut::from(&[0x82u8, 0x81, 1, 2][..]);
571 assert!(is_none(&Parser::parse(&mut buf, true, 1024)));
572
573 assert!(Parser::parse_close_payload(&[]).unwrap().is_none());
574 let reason = Parser::parse_close_payload(&[0x03, 0xe8, b'o', b'k'])
575 .unwrap()
576 .unwrap();
577 assert_eq!(reason.code, CloseCode::Normal);
578 assert_eq!(reason.description.as_deref(), Some("ok"));
579 }
580
581 #[test]
582 fn test_large_frame_roundtrip() {
583 let payload = Bytes::from(vec![7u8; 70_000]);
584 let mut buf = BytePages::default();
585 Parser::write_message(&mut buf, payload.clone(), OpCode::Binary, true, false).unwrap();
586
587 let mut buf = BytesMut::from(&Bytes::from(buf)[..]);
588 assert_eq!(&buf[..2], &[0x82, 127]);
589 let frame = extract(Parser::parse(&mut buf, false, 100_000));
590 assert!(frame.finished);
591 assert_eq!(frame.opcode, OpCode::Binary);
592 assert_eq!(frame.payload, payload);
593 }
594}