Skip to main content

tungstenite/protocol/frame/
mod.rs

1//! Utilities to work with raw WebSocket frames.
2
3pub mod coding;
4
5#[allow(clippy::module_inception)]
6mod frame;
7mod mask;
8mod utf8;
9
10pub use self::{
11    frame::{CloseFrame, Frame, FrameHeader},
12    utf8::Utf8Bytes,
13};
14
15use crate::{
16    error::{CapacityError, Error, ProtocolError, Result},
17    protocol::frame::mask::apply_mask,
18    Message,
19};
20use bytes::BytesMut;
21use log::*;
22use std::io::{self, Cursor, Error as IoError, ErrorKind as IoErrorKind, Read, Write};
23
24/// Read buffer size used for `FrameSocket`.
25const READ_BUF_LEN: usize = 128 * 1024;
26
27/// A reader and writer for WebSocket frames.
28#[derive(Debug)]
29pub struct FrameSocket<Stream> {
30    /// The underlying network stream.
31    stream: Stream,
32    /// Codec for reading/writing frames.
33    codec: FrameCodec,
34}
35
36impl<Stream> FrameSocket<Stream> {
37    /// Create a new frame socket.
38    pub fn new(stream: Stream) -> Self {
39        FrameSocket { stream, codec: FrameCodec::new(READ_BUF_LEN) }
40    }
41
42    /// Create a new frame socket from partially read data.
43    pub fn from_partially_read(stream: Stream, part: Vec<u8>) -> Self {
44        FrameSocket { stream, codec: FrameCodec::from_partially_read(part, READ_BUF_LEN) }
45    }
46
47    /// Extract a stream from the socket.
48    pub fn into_inner(self) -> (Stream, BytesMut) {
49        (self.stream, self.codec.in_buffer)
50    }
51
52    /// Returns a shared reference to the inner stream.
53    pub fn get_ref(&self) -> &Stream {
54        &self.stream
55    }
56
57    /// Returns a mutable reference to the inner stream.
58    pub fn get_mut(&mut self) -> &mut Stream {
59        &mut self.stream
60    }
61}
62
63impl<Stream> FrameSocket<Stream>
64where
65    Stream: Read,
66{
67    /// Read a frame from stream.
68    pub fn read(&mut self, max_size: Option<usize>) -> Result<Option<Frame>> {
69        self.codec.read_frame(&mut self.stream, max_size, false, true)
70    }
71}
72
73impl<Stream> FrameSocket<Stream>
74where
75    Stream: Write,
76{
77    /// Writes and immediately flushes a frame.
78    /// Equivalent to calling [`write`](Self::write) then [`flush`](Self::flush).
79    pub fn send(&mut self, frame: Frame) -> Result<()> {
80        self.write(frame)?;
81        self.flush()
82    }
83
84    /// Write a frame to stream.
85    ///
86    /// A subsequent call should be made to [`flush`](Self::flush) to flush writes.
87    ///
88    /// This function guarantees that the frame is queued unless [`Error::WriteBufferFull`]
89    /// is returned.
90    /// In order to handle WouldBlock or Incomplete, call [`flush`](Self::flush) afterwards.
91    pub fn write(&mut self, frame: Frame) -> Result<()> {
92        self.codec.buffer_frame(&mut self.stream, frame)
93    }
94
95    /// Flush writes.
96    pub fn flush(&mut self) -> Result<()> {
97        self.codec.write_out_buffer(&mut self.stream)?;
98        Ok(self.stream.flush()?)
99    }
100}
101
102/// A codec for WebSocket frames.
103#[derive(Debug)]
104pub(super) struct FrameCodec {
105    /// Buffer to read data from the stream.
106    in_buffer: BytesMut,
107    in_buf_max_read: usize,
108    /// Buffer to send packets to the network.
109    out_buffer: Vec<u8>,
110    /// Capacity limit for `out_buffer`.
111    max_out_buffer_len: usize,
112    /// Buffer target length to reach before writing to the stream
113    /// on calls to `buffer_frame`.
114    ///
115    /// Setting this to non-zero will buffer small writes from hitting
116    /// the stream.
117    out_buffer_write_len: usize,
118    /// Header and remaining size of the incoming packet being processed.
119    header: Option<(FrameHeader, u64)>,
120}
121
122impl FrameCodec {
123    /// Create a new frame codec.
124    pub(super) fn new(in_buf_len: usize) -> Self {
125        Self {
126            in_buffer: BytesMut::with_capacity(in_buf_len),
127            in_buf_max_read: in_buf_len.max(FrameHeader::MAX_SIZE),
128            out_buffer: <_>::default(),
129            max_out_buffer_len: usize::MAX,
130            out_buffer_write_len: 0,
131            header: None,
132        }
133    }
134
135    /// Create a new frame codec from partially read data.
136    pub(super) fn from_partially_read(part: Vec<u8>, min_in_buf_len: usize) -> Self {
137        let mut in_buffer = BytesMut::from_iter(part);
138        in_buffer.reserve(min_in_buf_len.saturating_sub(in_buffer.len()));
139        Self {
140            in_buffer,
141            in_buf_max_read: min_in_buf_len.max(FrameHeader::MAX_SIZE),
142            out_buffer: <_>::default(),
143            max_out_buffer_len: usize::MAX,
144            out_buffer_write_len: 0,
145            header: None,
146        }
147    }
148
149    /// Sets a maximum size for the out buffer.
150    pub(super) fn set_max_out_buffer_len(&mut self, max: usize) {
151        self.max_out_buffer_len = max;
152    }
153
154    /// Sets [`Self::buffer_frame`] buffer target length to reach before
155    /// writing to the stream.
156    pub(super) fn set_out_buffer_write_len(&mut self, len: usize) {
157        self.out_buffer_write_len = len;
158    }
159
160    /// Read a frame from the provided stream.
161    pub(super) fn read_frame(
162        &mut self,
163        stream: &mut impl Read,
164        max_size: Option<usize>,
165        unmask: bool,
166        accept_unmasked: bool,
167    ) -> Result<Option<Frame>> {
168        let max_size = max_size.unwrap_or_else(usize::max_value);
169
170        let mut payload = loop {
171            if self.header.is_none() {
172                let mut cursor = Cursor::new(&mut self.in_buffer);
173                self.header = FrameHeader::parse(&mut cursor)?;
174                let advanced = cursor.position();
175                bytes::Buf::advance(&mut self.in_buffer, advanced as _);
176
177                if let Some((_, len)) = &self.header {
178                    // Enforce frame size limit early. Compare in u64 before
179                    // narrowing so a length above usize::MAX can't wrap on
180                    // 32-bit/wasm32 targets and slip past the check.
181                    if *len > max_size as u64 {
182                        return Err(Error::Capacity(CapacityError::MessageTooLong {
183                            size: *len as usize,
184                            max_size,
185                        }));
186                    }
187                    let len = *len as usize;
188
189                    // Reserve full message length only once, even for multiple
190                    // loops or if WouldBlock errors cause multiple fn calls.
191                    self.in_buffer.reserve(len);
192                } else {
193                    self.in_buffer.reserve(FrameHeader::MAX_SIZE);
194                }
195            }
196
197            if let Some((_, len)) = &self.header {
198                let len = *len as usize;
199                if len <= self.in_buffer.len() {
200                    break self.in_buffer.split_to(len);
201                }
202            }
203
204            // Not enough data in buffer.
205            if self.read_in(stream)? == 0 {
206                trace!("no frame received");
207                return Ok(None);
208            }
209        };
210
211        let (mut header, length) = self.header.take().expect("Bug: no frame header");
212        debug_assert_eq!(payload.len() as u64, length);
213
214        if unmask {
215            if let Some(mask) = header.mask.take() {
216                // A server MUST remove masking for data frames received from a client
217                // as described in Section 5.3. (RFC 6455)
218                apply_mask(&mut payload, mask);
219            } else if !accept_unmasked {
220                // The server MUST close the connection upon receiving a
221                // frame that is not masked. (RFC 6455)
222                // The only exception here is if the user explicitly accepts given
223                // stream by setting WebSocketConfig.accept_unmasked_frames to true
224                return Err(Error::Protocol(ProtocolError::UnmaskedFrameFromClient));
225            }
226        }
227
228        let frame = Frame::from_payload(header, payload.freeze());
229        trace!("received frame {frame}");
230        Ok(Some(frame))
231    }
232
233    /// Read into available `in_buffer` capacity.
234    fn read_in(&mut self, stream: &mut impl Read) -> io::Result<usize> {
235        let len = self.in_buffer.len();
236        debug_assert!(self.in_buffer.capacity() > len);
237        self.in_buffer.resize(self.in_buffer.capacity().min(len + self.in_buf_max_read), 0);
238        let size = stream.read(&mut self.in_buffer[len..]);
239        self.in_buffer.truncate(len + size.as_ref().copied().unwrap_or(0));
240        size
241    }
242
243    /// Writes a frame into the `out_buffer`.
244    /// If the out buffer size is over the `out_buffer_write_len` will also write
245    /// the out buffer into the provided `stream`.
246    ///
247    /// To ensure buffered frames are written call [`Self::write_out_buffer`].
248    ///
249    /// May write to the stream, will **not** flush.
250    pub(super) fn buffer_frame<Stream>(&mut self, stream: &mut Stream, frame: Frame) -> Result<()>
251    where
252        Stream: Write,
253    {
254        if frame.len() + self.out_buffer.len() > self.max_out_buffer_len {
255            return Err(Error::WriteBufferFull(Message::Frame(frame).into()));
256        }
257
258        trace!("writing frame {frame}");
259
260        self.out_buffer.reserve(frame.len());
261        frame.format_into_buf(&mut self.out_buffer).expect("Bug: can't write to vector");
262
263        if self.out_buffer.len() > self.out_buffer_write_len {
264            self.write_out_buffer(stream)
265        } else {
266            Ok(())
267        }
268    }
269
270    /// Writes the out_buffer to the provided stream.
271    ///
272    /// Does **not** flush.
273    pub(super) fn write_out_buffer<Stream>(&mut self, stream: &mut Stream) -> Result<()>
274    where
275        Stream: Write,
276    {
277        while !self.out_buffer.is_empty() {
278            let len = stream.write(&self.out_buffer)?;
279            if len == 0 {
280                // This is the same as "Connection reset by peer"
281                return Err(IoError::new(
282                    IoErrorKind::ConnectionReset,
283                    "Connection reset while sending",
284                )
285                .into());
286            }
287            self.out_buffer.drain(0..len);
288        }
289
290        Ok(())
291    }
292}
293
294#[cfg(test)]
295mod tests {
296
297    use crate::error::{CapacityError, Error};
298
299    use super::{Frame, FrameSocket};
300
301    use std::io::Cursor;
302
303    #[test]
304    fn read_frames() {
305        env_logger::init();
306
307        let raw = Cursor::new(vec![
308            0x82, 0x07, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07, 0x82, 0x03, 0x03, 0x02, 0x01,
309            0x99,
310        ]);
311        let mut sock = FrameSocket::new(raw);
312
313        assert_eq!(
314            sock.read(None).unwrap().unwrap().into_payload(),
315            &[0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07][..]
316        );
317        assert_eq!(sock.read(None).unwrap().unwrap().into_payload(), &[0x03, 0x02, 0x01][..]);
318        assert!(sock.read(None).unwrap().is_none());
319
320        let (_, rest) = sock.into_inner();
321        assert_eq!(rest, vec![0x99]);
322    }
323
324    #[test]
325    fn from_partially_read() {
326        let raw = Cursor::new(vec![0x02, 0x03, 0x04, 0x05, 0x06, 0x07]);
327        let mut sock = FrameSocket::from_partially_read(raw, vec![0x82, 0x07, 0x01]);
328        assert_eq!(
329            sock.read(None).unwrap().unwrap().into_payload(),
330            &[0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07][..]
331        );
332    }
333
334    #[test]
335    fn write_frames() {
336        let mut sock = FrameSocket::new(Vec::new());
337
338        let frame = Frame::ping(vec![0x04, 0x05]);
339        sock.send(frame).unwrap();
340
341        let frame = Frame::pong(vec![0x01]);
342        sock.send(frame).unwrap();
343
344        let (buf, _) = sock.into_inner();
345        assert_eq!(buf, vec![0x89, 0x02, 0x04, 0x05, 0x8a, 0x01, 0x01]);
346    }
347
348    #[test]
349    fn parse_overflow() {
350        let raw = Cursor::new(vec![
351            0x83, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0xff, 0x00, 0x00, 0x00, 0x00,
352        ]);
353        let mut sock = FrameSocket::new(raw);
354        let _ = sock.read(None); // should not crash
355    }
356
357    #[test]
358    fn size_limit_hit() {
359        let raw = Cursor::new(vec![0x82, 0x07, 0x01, 0x02, 0x03, 0x04, 0x05, 0x06, 0x07]);
360        let mut sock = FrameSocket::new(raw);
361        assert!(matches!(
362            sock.read(Some(5)),
363            Err(Error::Capacity(CapacityError::MessageTooLong { size: 7, max_size: 5 }))
364        ));
365    }
366
367    #[test]
368    #[cfg(target_pointer_width = "32")]
369    fn length_above_usize_max_rejected() {
370        // 64-bit payload length 0x1_0000_0005 does not fit a 32-bit usize; it
371        // must be rejected rather than wrapping to 5 when narrowed.
372        let raw = Cursor::new(vec![0x82, 0x7f, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x05]);
373        let mut sock = FrameSocket::new(raw);
374        assert!(matches!(
375            sock.read(None),
376            Err(Error::Capacity(CapacityError::MessageTooLong { size: 5, max_size: usize::MAX }))
377        ));
378    }
379}