use crate::{
Message,
protocol::{PerMessageDeflateConfig, Role, WebSocket, WebSocketConfig, error::ProtocolError},
};
use rama_core::telemetry::tracing;
use std::io::{self, Cursor, Read, Write};
use tracing_test::traced_test;
pin_project_lite::pin_project! {
struct WriteMoc<Stream> {
#[pin]
stream: Stream,
written_bytes: usize,
write_count: usize,
flush_count: usize,
}
}
impl<Stream> WriteMoc<Stream> {
fn new(stream: Stream) -> Self {
Self {
stream,
written_bytes: 0,
write_count: 0,
flush_count: 0,
}
}
}
impl<Stream: Read> Read for WriteMoc<Stream> {
fn read(&mut self, buf: &mut [u8]) -> io::Result<usize> {
self.stream.read(buf)
}
}
impl<Stream> Write for WriteMoc<Stream> {
fn write(&mut self, buf: &[u8]) -> io::Result<usize> {
let n = buf.len();
self.written_bytes += n;
self.write_count += 1;
Ok(n)
}
fn flush(&mut self) -> io::Result<()> {
self.flush_count += 1;
Ok(())
}
}
#[test]
fn receive_messages() {
let incoming = Cursor::new(vec![
0x89, 0x02, 0x01, 0x02, 0x8a, 0x01, 0x03, 0x01, 0x07, 0x48, 0x65, 0x6c, 0x6c, 0x6f, 0x2c,
0x20, 0x80, 0x06, 0x57, 0x6f, 0x72, 0x6c, 0x64, 0x21, 0x82, 0x03, 0x01, 0x02, 0x03,
]);
let mut socket = WebSocket::from_raw_socket(WriteMoc::new(incoming), Role::Client, None);
assert_eq!(socket.read().unwrap(), Message::Ping(vec![1, 2].into()));
assert_eq!(socket.read().unwrap(), Message::Pong(vec![3].into()));
assert_eq!(
socket.read().unwrap(),
Message::Text("Hello, World!".into())
);
assert_eq!(
socket.read().unwrap(),
Message::Binary(vec![0x01, 0x02, 0x03].into())
);
}
#[test]
#[traced_test]
fn receive_compressed_hello_text_msg_rfc7692_example_compression_1() {
let incoming = Cursor::new(vec![0xc1, 0x07, 0xf2, 0x48, 0xcd, 0xc9, 0xc9, 0x07, 0x00]);
let mut socket = WebSocket::from_raw_socket(
WriteMoc::new(incoming),
Role::Client,
Some(WebSocketConfig {
per_message_deflate: Some(PerMessageDeflateConfig::default()),
..Default::default()
}),
);
assert_eq!(socket.read().unwrap(), Message::Text("Hello".into()));
}
#[test]
#[traced_test]
fn receive_compressed_hello_text_msg_rfc7692_example_compression_2() {
#[rustfmt::skip]
let incoming = Cursor::new(vec![
0x41, 0x03, 0xf2, 0x48, 0xcd, 0x80, 0x04, 0xc9, 0xc9, 0x07, 0x00, ]);
let mut socket = WebSocket::from_raw_socket(
WriteMoc::new(incoming),
Role::Client,
Some(WebSocketConfig {
per_message_deflate: Some(PerMessageDeflateConfig::default()),
..Default::default()
}),
);
assert_eq!(socket.read().unwrap(), Message::Text("Hello".into()));
}
#[test]
fn size_limiting_text_fragmented() {
let incoming = Cursor::new(vec![
0x01, 0x07, 0x48, 0x65, 0x6c, 0x6c, 0x6f, 0x2c, 0x20, 0x80, 0x06, 0x57, 0x6f, 0x72, 0x6c,
0x64, 0x21,
]);
let limit = WebSocketConfig {
max_message_size: Some(10),
..WebSocketConfig::default()
};
let mut socket = WebSocket::from_raw_socket(WriteMoc::new(incoming), Role::Client, Some(limit));
assert!(matches!(
socket.read(),
Err(ProtocolError::MessageTooLong {
size: 13,
max_size: 10
})
));
}
#[test]
fn size_limiting_binary() {
let incoming = Cursor::new(vec![0x82, 0x03, 0x01, 0x02, 0x03]);
let limit = WebSocketConfig {
max_message_size: Some(2),
..WebSocketConfig::default()
};
let mut socket = WebSocket::from_raw_socket(WriteMoc::new(incoming), Role::Client, Some(limit));
assert!(matches!(
socket.read(),
Err(ProtocolError::MessageTooLong {
size: 3,
max_size: 2
})
));
}
#[test]
fn server_write_flush_behaviour() {
const SEND_ME_LEN: usize = 10;
const BATCH_ME_LEN: usize = 11;
const WRITE_BUFFER_SIZE: usize = 600;
let mut ws = WebSocket::from_raw_socket(
WriteMoc::new(Cursor::new(Vec::default())),
Role::Server,
Some(WebSocketConfig::default().with_write_buffer_size(WRITE_BUFFER_SIZE)),
);
assert_eq!(ws.get_ref().written_bytes, 0);
assert_eq!(ws.get_ref().write_count, 0);
assert_eq!(ws.get_ref().flush_count, 0);
ws.send(Message::Text("Send me!".into())).unwrap();
assert_eq!(ws.get_ref().written_bytes, SEND_ME_LEN);
assert_eq!(ws.get_ref().write_count, 1);
assert_eq!(ws.get_ref().flush_count, 1);
for msg in (0..100).map(|_| Message::Text("Batch me!".into())) {
ws.write(msg).unwrap();
}
assert_eq!(ws.get_ref().written_bytes, 55 * BATCH_ME_LEN + SEND_ME_LEN);
assert_eq!(ws.get_ref().write_count, 2);
assert_eq!(ws.get_ref().flush_count, 1);
ws.flush().unwrap();
assert_eq!(ws.get_ref().written_bytes, 100 * BATCH_ME_LEN + SEND_ME_LEN);
assert_eq!(ws.get_ref().write_count, 3);
assert_eq!(ws.get_ref().flush_count, 2);
}
#[cfg(feature = "compression")]
fn raw_deflate_compress(data: &[u8]) -> Vec<u8> {
use flate2::{Compress, Compression, FlushCompress, Status};
let mut c = Compress::new_with_window_bits(Compression::default(), false, 15);
let mut buf: Vec<u8> = Vec::with_capacity(data.len());
let before = c.total_in();
while c.total_in() - before < data.len() as u64 {
buf.reserve(128);
match c
.compress_vec(
&data[(c.total_in() - before) as usize..],
&mut buf,
FlushCompress::Sync,
)
.unwrap()
{
Status::BufError => buf.reserve(buf.len()),
Status::Ok | Status::StreamEnd => {}
}
}
while !buf.ends_with(&[0x00, 0x00, 0xff, 0xff]) {
buf.reserve(8);
match c.compress_vec(&[], &mut buf, FlushCompress::Sync).unwrap() {
Status::BufError => buf.reserve(buf.len()),
Status::Ok | Status::StreamEnd => {}
}
}
buf.truncate(buf.len() - 4); buf
}
#[cfg(feature = "compression")]
fn ws_frame_compressed_binary(payload: &[u8]) -> Vec<u8> {
assert!(payload.len() < 126, "test helper only handles short frames");
let mut frame = vec![
0xC2, payload.len() as u8,
];
frame.extend_from_slice(payload);
frame
}
#[cfg(feature = "compression")]
#[test]
fn deflate_bomb_single_frame_rejected() {
const PLAIN_LEN: usize = 500;
let compressed = raw_deflate_compress(&vec![b'A'; PLAIN_LEN]);
let max_size = compressed.len() + 50;
assert!(
max_size < PLAIN_LEN,
"test invariant: max_size must be less than decompressed size"
);
let incoming = Cursor::new(ws_frame_compressed_binary(&compressed));
let limit = WebSocketConfig {
per_message_deflate: Some(PerMessageDeflateConfig::default()),
max_message_size: Some(max_size),
..Default::default()
};
let mut socket = WebSocket::from_raw_socket(WriteMoc::new(incoming), Role::Client, Some(limit));
match socket.read() {
Err(ProtocolError::MessageTooLong {
size,
max_size: reported_max,
}) => {
assert_eq!(reported_max, max_size);
assert!(size > reported_max);
assert!(size < PLAIN_LEN);
}
result => panic!("expected compressed payload to exceed decoded size limit: {result:?}"),
}
}
#[cfg(feature = "compression")]
#[test]
fn deflate_bomb_fragmented_rejected() {
const PLAIN_LEN: usize = 500;
let compressed = raw_deflate_compress(&vec![b'A'; PLAIN_LEN]);
let max_size = compressed.len() + 50;
assert!(
max_size < PLAIN_LEN,
"test invariant: max_size must be less than decompressed size"
);
let mid = compressed.len() / 2;
let (frag1, frag2) = compressed.split_at(mid);
let mut frame_bytes = Vec::new();
frame_bytes.push(0x42);
frame_bytes.push(frag1.len() as u8);
frame_bytes.extend_from_slice(frag1);
frame_bytes.push(0x80);
frame_bytes.push(frag2.len() as u8);
frame_bytes.extend_from_slice(frag2);
let incoming = Cursor::new(frame_bytes);
let limit = WebSocketConfig {
per_message_deflate: Some(PerMessageDeflateConfig::default()),
max_message_size: Some(max_size),
..Default::default()
};
let mut socket = WebSocket::from_raw_socket(WriteMoc::new(incoming), Role::Client, Some(limit));
match socket.read() {
Err(ProtocolError::MessageTooLong {
size,
max_size: reported_max,
}) => {
assert_eq!(reported_max, max_size);
assert!(size > reported_max);
assert!(size < PLAIN_LEN);
}
result => panic!("expected compressed payload to exceed decoded size limit: {result:?}"),
}
}