1pub 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
24const READ_BUF_LEN: usize = 128 * 1024;
26
27#[derive(Debug)]
29pub struct FrameSocket<Stream> {
30 stream: Stream,
32 codec: FrameCodec,
34}
35
36impl<Stream> FrameSocket<Stream> {
37 pub fn new(stream: Stream) -> Self {
39 FrameSocket { stream, codec: FrameCodec::new(READ_BUF_LEN) }
40 }
41
42 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 pub fn into_inner(self) -> (Stream, BytesMut) {
49 (self.stream, self.codec.in_buffer)
50 }
51
52 pub fn get_ref(&self) -> &Stream {
54 &self.stream
55 }
56
57 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 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 pub fn send(&mut self, frame: Frame) -> Result<()> {
80 self.write(frame)?;
81 self.flush()
82 }
83
84 pub fn write(&mut self, frame: Frame) -> Result<()> {
92 self.codec.buffer_frame(&mut self.stream, frame)
93 }
94
95 pub fn flush(&mut self) -> Result<()> {
97 self.codec.write_out_buffer(&mut self.stream)?;
98 Ok(self.stream.flush()?)
99 }
100}
101
102#[derive(Debug)]
104pub(super) struct FrameCodec {
105 in_buffer: BytesMut,
107 in_buf_max_read: usize,
108 out_buffer: Vec<u8>,
110 max_out_buffer_len: usize,
112 out_buffer_write_len: usize,
118 header: Option<(FrameHeader, u64)>,
120}
121
122impl FrameCodec {
123 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 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 pub(super) fn set_max_out_buffer_len(&mut self, max: usize) {
151 self.max_out_buffer_len = max;
152 }
153
154 pub(super) fn set_out_buffer_write_len(&mut self, len: usize) {
157 self.out_buffer_write_len = len;
158 }
159
160 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 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 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 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 apply_mask(&mut payload, mask);
219 } else if !accept_unmasked {
220 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 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 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 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 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); }
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 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}