Skip to main content

ntex/ws/
frame.rs

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/// WebSocket frame parser.
8#[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        // check masking
33        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        // Op code
41        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        // check for max allowed size
86        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    /// Parses one WebSocket frame from `src`.
106    ///
107    /// `server` selects the expected masking direction: server-side parsing
108    /// requires masked frames, while client-side parsing rejects them.
109    /// `max_size` limits the frame payload size.
110    ///
111    /// Returns the final-fragment flag, opcode, and optional payload when a
112    /// complete frame is available. Returns [`None`] without consuming a
113    /// partial frame.
114    pub fn parse(
115        src: &mut BytesMut,
116        server: bool,
117        max_size: usize,
118    ) -> Result<Option<(bool, OpCode, Option<Bytes>)>, ProtocolError> {
119        // try to parse ws frame metadata
120        let Some((idx, finished, opcode, length, mask)) =
121            Parser::parse_metadata(src, server, max_size)?
122        else {
123            return Ok(None);
124        };
125
126        // not enough data
127        if src.len() < idx + length {
128            return Ok(None);
129        }
130
131        // remove prefix
132        src.advance_to(idx);
133
134        // no need for body
135        if length == 0 {
136            return Ok(Some((finished, opcode, None)));
137        }
138
139        // unmask
140        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    /// Parses a close-frame payload.
148    ///
149    /// Returns [`None`] for an empty payload.
150    ///
151    /// # Errors
152    ///
153    /// Returns an error for a one-byte payload, an invalid close status code,
154    /// or a description that is not valid UTF-8.
155    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    /// Encodes a WebSocket frame into `dst`.
182    ///
183    /// `fin` controls the final-fragment bit and `mask` controls whether a new
184    /// random masking key is applied.
185    ///
186    /// # Errors
187    ///
188    /// Returns an error if a control frame is fragmented or has a payload
189    /// larger than 125 bytes.
190    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    /// Encodes a final close control frame into `dst`.
241    ///
242    /// # Errors
243    ///
244    /// Returns an error if the close code cannot be sent or the encoded close
245    /// payload would exceed 125 bytes.
246    #[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        // 64-bit length over the max size
563        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        // masking key is not received yet
570        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}