Skip to main content

actix_http/ws/
frame.rs

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/// A struct representing a WebSocket frame.
13#[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        // RSV1, RSV2, and RSV3 must be zero unless a negotiated extension defines them.
33        if first & 0b0111_0000 != 0 {
34            // TODO(semver-major): use InvalidReservedBits
35            return Err(ProtocolError::Io(io::Error::new(
36                io::ErrorKind::InvalidData,
37                "Received a frame with non-zero reserved bits",
38            )));
39        }
40
41        // check masking
42        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        // Op code
50        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    /// Parse the input stream into a frame.
95    pub fn parse(
96        src: &mut BytesMut,
97        server: bool,
98        max_size: usize,
99    ) -> Result<Option<(bool, OpCode, Option<BytesMut>)>, ProtocolError> {
100        // try to parse ws frame metadata
101        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        // not enough data
112        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        // remove prefix
126        src.advance(idx);
127
128        // check for max allowed size
129        if length > max_size {
130            // drop the payload
131            src.advance(length);
132            return Err(ProtocolError::Overflow);
133        }
134
135        // no need for body
136        if length == 0 {
137            return Ok(Some((finished, opcode, None)));
138        }
139
140        let mut data = src.split_to(length);
141
142        // control frames must have length <= 125
143        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        // unmask
155        if let Some(mask) = mask {
156            apply_mask(&mut data, mask);
157        }
158
159        Ok(Some((finished, opcode, Some(data))))
160    }
161
162    /// Parse the payload of a close frame.
163    ///
164    /// This method preserves the historical behavior of accepting unknown status codes and
165    /// replacing invalid UTF-8 in the reason. Use [`Parser::try_parse_close_payload`] when the
166    /// payload must be validated.
167    #[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    /// Parse and validate the payload of a close frame.
187    ///
188    /// # Validation
189    ///
190    /// A close payload is either empty or contains a two-byte status code followed by an optional
191    /// UTF-8 reason. This follows [RFC 6455 §5.5.1], with status code rules from [RFC 6455 §7.4]
192    /// and UTF-8 error handling from [RFC 6455 §8.1].
193    ///
194    /// [RFC 6455 §5.5.1]: https://datatracker.ietf.org/doc/html/rfc6455#section-5.5.1
195    /// [RFC 6455 §7.4]: https://datatracker.ietf.org/doc/html/rfc6455#section-7.4
196    /// [RFC 6455 §8.1]: https://datatracker.ietf.org/doc/html/rfc6455#section-8.1
197    pub fn try_parse_close_payload(payload: &[u8]) -> Result<Option<CloseReason>, ProtocolError> {
198        // RFC 6455 §5.5.1 requires a two-byte status code when a close payload is not empty.
199        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            // RFC 6455 §7.4 reserves 1000-2999 for protocol-defined codes. Reject reserved and
211            // undefined values while allowing registered codes (3000-3999) and private-use codes
212            // (4000-4999).
213            if !matches!(raw_code, 1000..=1003 | 1007..=1014 | 3000..=4999) {
214                // TODO(semver-major): use this instead
215                // return Err(ProtocolError::InvalidCloseCode(raw_code));
216                return Err(ProtocolError::BadOpCode);
217            }
218
219            if payload.len() > 2 {
220                // RFC 6455 §5.5.1 defines the remaining bytes as a UTF-8 reason. RFC 6455 §8.1
221                // requires the connection to fail when data interpreted as UTF-8 is invalid.
222                str::from_utf8(&payload[2..]).map_err(|_| {
223                    // TODO(semver-major): use this instead
224                    // ProtocolError::InvalidCloseReason
225                    ProtocolError::BadOpCode
226                })?;
227            }
228        }
229
230        #[expect(deprecated)]
231        Ok(Self::parse_close_payload(payload))
232    }
233
234    /// Generate binary representation
235    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    /// Create a new Close control frame.
280    #[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        // https://github.com/actix/actix-web/issues/1579
409
410        // Final, masked text frame with RSV1 set and no negotiated extension.
411        let mut buf = BytesMut::from(
412            &[
413                0b1100_0001u8, // FIN + RSV1 + text
414                0b1000_0001u8, // MASK + payload length
415                0,             // masking key byte 1
416                0,             // masking key byte 2
417                0,             // masking key byte 3
418                0,             // masking key byte 4
419                b'a',          // payload
420            ][..],
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        // The first text frame with length == 2, payload doesn't matter.
444        buf.extend([0b0000_0001u8, 0b0000_0010u8, 0b0000_0000u8, 0b0000_0000u8]);
445        // Next binary frame with length == 2 and payload == `[0x1111_1111u8, 0x1111_1111u8]`.
446        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}