Skip to main content

ntex/ws/
codec.rs

1use std::cell::Cell;
2
3use crate::codec::{Decoder, Encoder};
4use crate::util::{BytePage, BytePages, ByteString, Bytes, BytesMut};
5
6use super::error::ProtocolError;
7use super::frame::Parser;
8use super::proto::{CloseReason, OpCode};
9
10/// An outgoing WebSocket message.
11#[derive(Debug, Clone, PartialEq, Eq)]
12pub enum Message {
13    /// UTF-8 text message.
14    Text(ByteString),
15    /// Binary message.
16    Binary(Bytes),
17    /// Fragment of a text or binary message.
18    Continuation(Item),
19    /// Ping control message.
20    Ping(Bytes),
21    /// Pong control message.
22    Pong(Bytes),
23    /// Close control message with an optional reason.
24    Close(Option<CloseReason>),
25}
26
27/// A decoded WebSocket frame.
28#[derive(Debug, Clone, PartialEq, Eq)]
29pub enum Frame {
30    /// Text frame.
31    ///
32    /// The codec does not validate that the payload is UTF-8.
33    Text(Bytes),
34    /// Binary frame.
35    Binary(Bytes),
36    /// Fragmented text or binary frame.
37    Continuation(Item),
38    /// Ping control frame.
39    Ping(Bytes),
40    /// Pong control frame.
41    Pong(Bytes),
42    /// Close control frame with an optional reason.
43    Close(Option<CloseReason>),
44}
45
46/// A fragment in a WebSocket continuation sequence.
47#[derive(Debug, Clone, PartialEq, Eq)]
48pub enum Item {
49    /// First fragment of a text message.
50    FirstText(Bytes),
51    /// First fragment of a binary message.
52    FirstBinary(Bytes),
53    /// Intermediate fragment.
54    Continue(Bytes),
55    /// Final fragment.
56    Last(Bytes),
57}
58
59#[derive(Debug, Clone)]
60/// Encoder and decoder for WebSocket frames.
61pub struct Codec {
62    flags: Cell<Flags>,
63    max_size: usize,
64}
65
66bitflags::bitflags! {
67    #[derive(Copy, Clone, Debug, Eq, PartialEq, Ord, PartialOrd, Hash)]
68    struct Flags: u8 {
69        const SERVER         = 0b0000_0001;
70        const R_CONTINUATION = 0b0000_0010;
71        const W_CONTINUATION = 0b0000_0100;
72        const CLOSED         = 0b0000_1000;
73        const R_CLOSED       = 0b0001_0000;
74    }
75}
76
77impl Codec {
78    /// Creates a codec in server mode with a 64 KiB frame-size limit.
79    #[must_use]
80    pub fn new() -> Codec {
81        Codec {
82            max_size: 65_536,
83            flags: Cell::new(Flags::SERVER),
84        }
85    }
86
87    /// Sets the maximum accepted frame payload size.
88    ///
89    /// The default is 64 KiB.
90    #[must_use]
91    pub fn max_size(mut self, size: usize) -> Self {
92        self.max_size = size;
93        self
94    }
95
96    /// Configures the codec for client-side masking rules.
97    ///
98    /// Client mode masks encoded frames and rejects masked incoming frames.
99    /// By default, the codec uses server-side masking rules.
100    #[must_use]
101    pub fn set_client_mode(self) -> Self {
102        self.remove_flags(Flags::SERVER);
103        self
104    }
105
106    /// Returns `true` after this codec has encoded a close message.
107    pub fn is_closed(&self) -> bool {
108        self.flags().contains(Flags::CLOSED)
109    }
110
111    fn flags(&self) -> Flags {
112        self.flags.get()
113    }
114
115    fn insert_flags(&self, f: Flags) {
116        self.flags.set(self.flags() | f);
117    }
118
119    fn remove_flags(&self, f: Flags) {
120        self.flags.set(self.flags() - f);
121    }
122
123    /// Encodes `page` as a final binary frame.
124    ///
125    /// # Errors
126    ///
127    /// Returns [`ProtocolError::Closed`] if a close message has already been
128    /// encoded, or [`ProtocolError::ContinuationStarted`] if a fragmented
129    /// message is in progress.
130    pub fn encode_page(&self, page: BytePage, dst: &mut BytePages) -> Result<(), ProtocolError> {
131        if self.is_closed() {
132            return Err(ProtocolError::Closed);
133        }
134        if self.flags().contains(Flags::W_CONTINUATION) {
135            return Err(ProtocolError::ContinuationStarted);
136        }
137        Parser::write_message(
138            dst,
139            page,
140            OpCode::Binary,
141            true,
142            !self.flags().contains(Flags::SERVER),
143        )
144        .expect("binary frames are always valid");
145        Ok(())
146    }
147}
148
149impl Default for Codec {
150    fn default() -> Self {
151        Self::new()
152    }
153}
154
155impl Encoder for Codec {
156    type Item = Message;
157    type Error = ProtocolError;
158
159    fn encode(&self, item: Message, dst: &mut BytePages) -> Result<(), Self::Error> {
160        if self.is_closed() {
161            return Err(ProtocolError::Closed);
162        }
163
164        match item {
165            Message::Text(txt) => {
166                if self.flags().contains(Flags::W_CONTINUATION) {
167                    return Err(ProtocolError::ContinuationStarted);
168                }
169                Parser::write_message(
170                    dst,
171                    txt,
172                    OpCode::Text,
173                    true,
174                    !self.flags().contains(Flags::SERVER),
175                )?;
176            }
177            Message::Binary(bin) => {
178                if self.flags().contains(Flags::W_CONTINUATION) {
179                    return Err(ProtocolError::ContinuationStarted);
180                }
181                Parser::write_message(
182                    dst,
183                    bin,
184                    OpCode::Binary,
185                    true,
186                    !self.flags().contains(Flags::SERVER),
187                )?;
188            }
189            Message::Ping(txt) => Parser::write_message(
190                dst,
191                txt,
192                OpCode::Ping,
193                true,
194                !self.flags().contains(Flags::SERVER),
195            )?,
196            Message::Pong(txt) => Parser::write_message(
197                dst,
198                txt,
199                OpCode::Pong,
200                true,
201                !self.flags().contains(Flags::SERVER),
202            )?,
203            Message::Close(reason) => {
204                Parser::write_close(dst, reason, !self.flags().contains(Flags::SERVER))?;
205                self.insert_flags(Flags::CLOSED);
206            }
207            Message::Continuation(cont) => match cont {
208                Item::FirstText(data) => {
209                    if self.flags().contains(Flags::W_CONTINUATION) {
210                        return Err(ProtocolError::ContinuationStarted);
211                    }
212                    self.insert_flags(Flags::W_CONTINUATION);
213                    Parser::write_message(
214                        dst,
215                        data,
216                        OpCode::Text,
217                        false,
218                        !self.flags().contains(Flags::SERVER),
219                    )?;
220                }
221                Item::FirstBinary(data) => {
222                    if self.flags().contains(Flags::W_CONTINUATION) {
223                        return Err(ProtocolError::ContinuationStarted);
224                    }
225                    self.insert_flags(Flags::W_CONTINUATION);
226                    Parser::write_message(
227                        dst,
228                        data,
229                        OpCode::Binary,
230                        false,
231                        !self.flags().contains(Flags::SERVER),
232                    )?;
233                }
234                Item::Continue(data) => {
235                    if self.flags().contains(Flags::W_CONTINUATION) {
236                        Parser::write_message(
237                            dst,
238                            data,
239                            OpCode::Continue,
240                            false,
241                            !self.flags().contains(Flags::SERVER),
242                        )?;
243                    } else {
244                        return Err(ProtocolError::ContinuationNotStarted);
245                    }
246                }
247                Item::Last(data) => {
248                    if self.flags().contains(Flags::W_CONTINUATION) {
249                        self.remove_flags(Flags::W_CONTINUATION);
250                        Parser::write_message(
251                            dst,
252                            data,
253                            OpCode::Continue,
254                            true,
255                            !self.flags().contains(Flags::SERVER),
256                        )?;
257                    } else {
258                        return Err(ProtocolError::ContinuationNotStarted);
259                    }
260                }
261            },
262        }
263        Ok(())
264    }
265}
266
267impl Decoder for Codec {
268    type Item = Frame;
269    type Error = ProtocolError;
270
271    fn decode(&self, src: &mut BytesMut) -> Result<Option<Self::Item>, Self::Error> {
272        // the peer must not send anything after its close frame, discard it
273        if self.flags().contains(Flags::R_CLOSED) {
274            src.clear();
275            return Ok(None);
276        }
277
278        match Parser::parse(src, self.flags().contains(Flags::SERVER), self.max_size) {
279            Ok(Some((finished, opcode, payload))) => {
280                // handle continuation
281                if finished {
282                    match opcode {
283                        OpCode::Continue => {
284                            if self.flags().contains(Flags::R_CONTINUATION) {
285                                self.remove_flags(Flags::R_CONTINUATION);
286                                let payload = payload.unwrap_or_default();
287                                Ok(Some(Frame::Continuation(Item::Last(payload))))
288                            } else {
289                                Err(ProtocolError::ContinuationNotStarted)
290                            }
291                        }
292                        OpCode::Close => {
293                            let reason = if let Some(pl) = payload {
294                                Parser::parse_close_payload(&pl)?
295                            } else {
296                                None
297                            };
298                            self.insert_flags(Flags::R_CLOSED);
299                            Ok(Some(Frame::Close(reason)))
300                        }
301                        OpCode::Ping => Ok(Some(Frame::Ping(payload.unwrap_or_default()))),
302                        OpCode::Pong => Ok(Some(Frame::Pong(payload.unwrap_or_default()))),
303                        OpCode::Binary => {
304                            if self.flags().contains(Flags::R_CONTINUATION) {
305                                Err(ProtocolError::ContinuationStarted)
306                            } else {
307                                Ok(Some(Frame::Binary(payload.unwrap_or_else(Bytes::new))))
308                            }
309                        }
310                        OpCode::Text => {
311                            if self.flags().contains(Flags::R_CONTINUATION) {
312                                Err(ProtocolError::ContinuationStarted)
313                            } else {
314                                Ok(Some(Frame::Text(payload.unwrap_or_else(Bytes::new))))
315                            }
316                        }
317                    }
318                } else {
319                    match opcode {
320                        OpCode::Continue => {
321                            if self.flags().contains(Flags::R_CONTINUATION) {
322                                Ok(Some(Frame::Continuation(Item::Continue(
323                                    payload.unwrap_or_else(Bytes::new),
324                                ))))
325                            } else {
326                                Err(ProtocolError::ContinuationNotStarted)
327                            }
328                        }
329                        OpCode::Binary => {
330                            if self.flags().contains(Flags::R_CONTINUATION) {
331                                Err(ProtocolError::ContinuationStarted)
332                            } else {
333                                self.insert_flags(Flags::R_CONTINUATION);
334                                Ok(Some(Frame::Continuation(Item::FirstBinary(
335                                    payload.unwrap_or_else(Bytes::new),
336                                ))))
337                            }
338                        }
339                        OpCode::Text => {
340                            if self.flags().contains(Flags::R_CONTINUATION) {
341                                Err(ProtocolError::ContinuationStarted)
342                            } else {
343                                self.insert_flags(Flags::R_CONTINUATION);
344                                Ok(Some(Frame::Continuation(Item::FirstText(
345                                    payload.unwrap_or_else(Bytes::new),
346                                ))))
347                            }
348                        }
349                        // rejected by the parser, kept for exhaustiveness
350                        OpCode::Ping | OpCode::Pong | OpCode::Close => {
351                            Err(ProtocolError::FragmentedControlFrame(opcode))
352                        }
353                    }
354                }
355            }
356            Ok(None) => Ok(None),
357            Err(e) => Err(e),
358        }
359    }
360}
361
362#[cfg(test)]
363mod tests {
364    use super::*;
365    use crate::ws::CloseCode;
366
367    #[test]
368    fn text_payload_is_not_validated() {
369        let codec = Codec::new().set_client_mode();
370        let mut frame = BytesMut::from(&[0x81, 0x01, 0xff][..]);
371        assert!(matches!(
372            codec.decode(&mut frame),
373            Ok(Some(Frame::Text(data))) if data == Bytes::from_static(&[0xff])
374        ));
375    }
376
377    #[test]
378    fn input_after_close_is_discarded() {
379        let codec = Codec::new().set_client_mode();
380        // close frame, followed by a text frame
381        let mut src = BytesMut::from(&[0x88, 0x00, 0x81, 0x01, b'a'][..]);
382        assert!(matches!(
383            codec.decode(&mut src),
384            Ok(Some(Frame::Close(None)))
385        ));
386        assert!(matches!(codec.decode(&mut src), Ok(None)));
387        assert!(src.is_empty());
388
389        src.extend_from_slice(&[0x81, 0x01, b'b']);
390        assert!(matches!(codec.decode(&mut src), Ok(None)));
391        assert!(src.is_empty());
392    }
393
394    #[test]
395    fn accepts_extension_close_code() {
396        // servers should not send 1010, but receivers accept it
397        let codec = Codec::new().set_client_mode();
398        let mut close = BytesMut::from(&[0x88, 0x02, 0x03, 0xf2][..]);
399        assert!(matches!(
400            codec.decode(&mut close),
401            Ok(Some(Frame::Close(Some(CloseReason {
402                code: CloseCode::Extension,
403                ..
404            }))))
405        ));
406
407        // servers still cannot send it
408        let codec = Codec::new();
409        let mut dst = BytePages::default();
410        assert!(matches!(
411            codec.encode(Message::Close(Some(CloseCode::Extension.into())), &mut dst),
412            Err(ProtocolError::InvalidCloseCode(1010))
413        ));
414    }
415
416    #[test]
417    fn validates_outgoing_continuations() {
418        let codec = Codec::new();
419        let mut dst = BytePages::default();
420        codec
421            .encode(
422                Message::Continuation(Item::FirstBinary(Bytes::new())),
423                &mut dst,
424            )
425            .unwrap();
426        assert!(matches!(
427            codec.encode(Message::Text("text".into()), &mut dst),
428            Err(ProtocolError::ContinuationStarted)
429        ));
430        assert!(matches!(
431            codec.encode_page(BytePage::from(Bytes::new()), &mut dst),
432            Err(ProtocolError::ContinuationStarted)
433        ));
434
435        codec
436            .encode(Message::Continuation(Item::Last(Bytes::new())), &mut dst)
437            .unwrap();
438        codec
439            .encode_page(BytePage::from(Bytes::new()), &mut dst)
440            .unwrap();
441    }
442
443    #[test]
444    fn rejects_messages_after_close() {
445        let codec = Codec::new();
446        let mut dst = BytePages::default();
447        codec.encode(Message::Close(None), &mut dst).unwrap();
448
449        assert!(matches!(
450            codec.encode(Message::Text("text".into()), &mut dst),
451            Err(ProtocolError::Closed)
452        ));
453        assert!(matches!(
454            codec.encode_page(BytePage::from(Bytes::new()), &mut dst),
455            Err(ProtocolError::Closed)
456        ));
457    }
458
459    #[test]
460    fn encode_errors() {
461        let codec = Codec::new();
462        let mut dst = BytePages::default();
463        let big = Bytes::from(vec![0; 126]);
464        assert!(matches!(
465            codec.encode(Message::Ping(big.clone()), &mut dst),
466            Err(ProtocolError::InvalidLength(126))
467        ));
468        assert!(matches!(
469            codec.encode(Message::Pong(big), &mut dst),
470            Err(ProtocolError::InvalidLength(126))
471        ));
472
473        codec
474            .encode(Message::Continuation(Item::FirstText("a".into())), &mut dst)
475            .unwrap();
476        assert!(matches!(
477            codec.encode(Message::Binary("b".into()), &mut dst),
478            Err(ProtocolError::ContinuationStarted)
479        ));
480        assert!(matches!(
481            codec.encode(
482                Message::Continuation(Item::FirstBinary("b".into())),
483                &mut dst
484            ),
485            Err(ProtocolError::ContinuationStarted)
486        ));
487        codec
488            .encode(Message::Continuation(Item::Last("c".into())), &mut dst)
489            .unwrap();
490        assert!(matches!(
491            codec.encode(Message::Continuation(Item::Continue("d".into())), &mut dst),
492            Err(ProtocolError::ContinuationNotStarted)
493        ));
494    }
495
496    fn decode(codec: &Codec, frame: &[u8]) -> Result<Option<Frame>, ProtocolError> {
497        codec.decode(&mut BytesMut::from(frame))
498    }
499
500    #[test]
501    fn decode_continuation_errors() {
502        let codec = Codec::new().set_client_mode();
503        // continuation without a first frame
504        assert!(matches!(
505            decode(&codec, &[0x80, 0x01, b'a']),
506            Err(ProtocolError::ContinuationNotStarted)
507        ));
508        assert!(matches!(
509            decode(&codec, &[0x00, 0x01, b'a']),
510            Err(ProtocolError::ContinuationNotStarted)
511        ));
512
513        // new data frames while a fragmented message is in progress
514        assert!(matches!(
515            decode(&codec, &[0x02, 0x01, b'a']),
516            Ok(Some(Frame::Continuation(Item::FirstBinary(_))))
517        ));
518        for frame in [
519            &[0x82, 0x01, b'a'],
520            &[0x81, 0x01, b'a'],
521            &[0x02, 0x01, b'a'],
522            &[0x01, 0x01, b'a'],
523        ] {
524            assert!(matches!(
525                decode(&codec, frame),
526                Err(ProtocolError::ContinuationStarted)
527            ));
528        }
529        assert!(matches!(
530            decode(&codec, &[0x00, 0x01, b'b']),
531            Ok(Some(Frame::Continuation(Item::Continue(_))))
532        ));
533        assert!(matches!(
534            decode(&codec, &[0x80, 0x00]),
535            Ok(Some(Frame::Continuation(Item::Last(data)))) if data.is_empty()
536        ));
537    }
538}