Skip to content

Commit bc39c1f

Browse files
divybotlittledivy
andauthored
fix(fragment): cap fragmented message size (#140)
Co-authored-by: divybot <divybot@users.noreply.github.com> Co-authored-by: Divy Srivastava <me@littledivy.com>
1 parent 9cc13ea commit bc39c1f

2 files changed

Lines changed: 118 additions & 6 deletions

File tree

src/fragment.rs

Lines changed: 37 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -27,15 +27,15 @@ use tokio::io::AsyncRead;
2727
use tokio::io::AsyncWrite;
2828

2929
pub enum Fragment {
30-
Text(Option<utf8::Incomplete>, Vec<u8>),
30+
Text(Option<utf8::Incomplete>, Vec<u8>, usize),
3131
Binary(Vec<u8>),
3232
}
3333

3434
impl Fragment {
3535
/// Returns the payload of the fragment.
3636
fn take_buffer(self) -> Vec<u8> {
3737
match self {
38-
Fragment::Text(_, buffer) => buffer,
38+
Fragment::Text(_, buffer, _) => buffer,
3939
Fragment::Binary(buffer) => buffer,
4040
}
4141
}
@@ -118,7 +118,10 @@ impl<'f, S> FragmentCollector<S> {
118118
if is_closed && frame.opcode != OpCode::Close {
119119
return Err(WebSocketError::ConnectionClosed);
120120
}
121-
if let Some(frame) = self.fragments.accumulate(frame)? {
121+
if let Some(frame) = self
122+
.fragments
123+
.accumulate(frame, self.read_half.max_message_size)?
124+
{
122125
return Ok(frame);
123126
}
124127
}
@@ -191,7 +194,10 @@ impl<'f, S> FragmentCollectorRead<S> {
191194
let Some(frame) = res? else {
192195
continue;
193196
};
194-
if let Some(frame) = self.fragments.accumulate(frame)? {
197+
if let Some(frame) = self
198+
.fragments
199+
.accumulate(frame, self.read_half.max_message_size)?
200+
{
195201
return Ok(frame);
196202
}
197203
}
@@ -215,6 +221,7 @@ impl Fragments {
215221
pub fn accumulate<'f>(
216222
&mut self,
217223
frame: Frame<'f>,
224+
max_message_size: usize,
218225
) -> Result<Option<Frame<'f>>, WebSocketError> {
219226
match frame.opcode {
220227
OpCode::Text | OpCode::Binary => {
@@ -224,15 +231,23 @@ impl Fragments {
224231
}
225232
return Ok(Some(Frame::new(true, frame.opcode, None, frame.payload)));
226233
} else {
234+
if frame.payload.len() >= max_message_size {
235+
return Err(WebSocketError::FrameTooLarge);
236+
}
227237
self.fragments = match frame.opcode {
228238
OpCode::Text => match utf8::decode(&frame.payload) {
229-
Ok(text) => Some(Fragment::Text(None, text.as_bytes().to_vec())),
239+
Ok(text) => Some(Fragment::Text(
240+
None,
241+
text.as_bytes().to_vec(),
242+
frame.payload.len(),
243+
)),
230244
Err(utf8::DecodeError::Incomplete {
231245
valid_prefix,
232246
incomplete_suffix,
233247
}) => Some(Fragment::Text(
234248
Some(incomplete_suffix),
235249
valid_prefix.as_bytes().to_vec(),
250+
frame.payload.len(),
236251
)),
237252
Err(utf8::DecodeError::Invalid { .. }) => {
238253
return Err(WebSocketError::InvalidUTF8);
@@ -248,7 +263,15 @@ impl Fragments {
248263
None => {
249264
return Err(WebSocketError::InvalidContinuationFrame);
250265
}
251-
Some(Fragment::Text(data, input)) => {
266+
Some(Fragment::Text(data, input, message_len)) => {
267+
let new_message_len = message_len
268+
.checked_add(frame.payload.len())
269+
.ok_or(WebSocketError::FrameTooLarge)?;
270+
if new_message_len >= max_message_size {
271+
return Err(WebSocketError::FrameTooLarge);
272+
}
273+
*message_len = new_message_len;
274+
252275
let mut tail = &frame.payload[..];
253276
if let Some(mut incomplete) = data.take() {
254277
if let Some((result, rest)) =
@@ -296,6 +319,14 @@ impl Fragments {
296319
}
297320
}
298321
Some(Fragment::Binary(data)) => {
322+
let message_len = data
323+
.len()
324+
.checked_add(frame.payload.len())
325+
.ok_or(WebSocketError::FrameTooLarge)?;
326+
if message_len >= max_message_size {
327+
return Err(WebSocketError::FrameTooLarge);
328+
}
329+
299330
data.extend_from_slice(&frame.payload);
300331
if frame.fin {
301332
return Ok(Some(Frame::new(

tests/fragment.rs

Lines changed: 81 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,81 @@
1+
use fastwebsockets::FragmentCollector;
2+
#[cfg(feature = "unstable-split")]
3+
use fastwebsockets::FragmentCollectorRead;
4+
use fastwebsockets::Frame;
5+
use fastwebsockets::OpCode;
6+
use fastwebsockets::Role;
7+
use fastwebsockets::WebSocket;
8+
use fastwebsockets::WebSocketError;
9+
use tokio::io::AsyncWriteExt;
10+
11+
fn encoded_frames(mut frames: Vec<Frame<'static>>) -> Vec<u8> {
12+
let mut out = Vec::new();
13+
let mut scratch = Vec::new();
14+
15+
for frame in &mut frames {
16+
out.extend_from_slice(frame.write(&mut scratch));
17+
}
18+
19+
out
20+
}
21+
22+
fn assert_frame_too_large<T>(result: Result<T, WebSocketError>) {
23+
assert!(matches!(result, Err(WebSocketError::FrameTooLarge)));
24+
}
25+
26+
#[tokio::test]
27+
async fn fragment_collector_rejects_aggregate_binary_over_limit() {
28+
let (mut peer, socket) = tokio::io::duplex(1024);
29+
let mut ws = WebSocket::after_handshake(socket, Role::Client);
30+
ws.set_max_message_size(9);
31+
let mut ws = FragmentCollector::new(ws);
32+
33+
let frames = encoded_frames(vec![
34+
Frame::new(false, OpCode::Binary, None, b"12345".to_vec().into()),
35+
Frame::new(true, OpCode::Continuation, None, b"67890".to_vec().into()),
36+
]);
37+
peer.write_all(&frames).await.unwrap();
38+
39+
assert_frame_too_large(ws.read_frame().await);
40+
}
41+
42+
#[tokio::test]
43+
async fn fragment_collector_rejects_aggregate_text_over_limit() {
44+
let (mut peer, socket) = tokio::io::duplex(1024);
45+
let mut ws = WebSocket::after_handshake(socket, Role::Client);
46+
ws.set_max_message_size(9);
47+
let mut ws = FragmentCollector::new(ws);
48+
49+
let frames = encoded_frames(vec![
50+
Frame::new(false, OpCode::Text, None, b"hello".to_vec().into()),
51+
Frame::new(true, OpCode::Continuation, None, b"world".to_vec().into()),
52+
]);
53+
peer.write_all(&frames).await.unwrap();
54+
55+
assert_frame_too_large(ws.read_frame().await);
56+
}
57+
58+
#[cfg(feature = "unstable-split")]
59+
#[tokio::test]
60+
async fn split_fragment_collector_rejects_aggregate_binary_over_limit() {
61+
let (mut peer, socket) = tokio::io::duplex(1024);
62+
let (read, _write) = tokio::io::split(socket);
63+
let (mut ws_read, _ws_write) = fastwebsockets::after_handshake_split(
64+
read,
65+
tokio::io::sink(),
66+
Role::Client,
67+
);
68+
ws_read.set_max_message_size(9);
69+
let mut ws = FragmentCollectorRead::new(ws_read);
70+
71+
let frames = encoded_frames(vec![
72+
Frame::new(false, OpCode::Binary, None, b"12345".to_vec().into()),
73+
Frame::new(true, OpCode::Continuation, None, b"67890".to_vec().into()),
74+
]);
75+
peer.write_all(&frames).await.unwrap();
76+
77+
assert_frame_too_large(
78+
ws.read_frame(&mut |_| async { Ok::<(), std::io::Error>(()) })
79+
.await,
80+
);
81+
}

0 commit comments

Comments
 (0)