use std::io::{self, Read, Write};
pub const MAX_FRAME_LEN: usize = 64 * 1024 * 1024;
pub fn header(len: usize) -> io::Result<[u8; 4]> {
u32::try_from(len)
.map(u32::to_be_bytes)
.map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "frame too large"))
}
pub fn payload_len(header: [u8; 4]) -> io::Result<usize> {
let len = u32::from_be_bytes(header) as usize;
if len > MAX_FRAME_LEN {
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!("frame of {len} bytes exceeds the {MAX_FRAME_LEN} byte limit"),
));
}
Ok(len)
}
pub fn write_frame<W: Write>(mut w: W, payload: &[u8]) -> io::Result<()> {
w.write_all(&header(payload.len())?)?;
w.write_all(payload)?;
w.flush()
}
pub fn read_frame<R: Read>(mut r: R) -> io::Result<Option<Vec<u8>>> {
let mut first = [0u8; 1];
if r.read(&mut first)? == 0 {
return Ok(None);
}
let mut rest = [0u8; 3];
r.read_exact(&mut rest)?;
let len = payload_len([first[0], rest[0], rest[1], rest[2]])?;
let mut payload = vec![0u8; len];
r.read_exact(&mut payload)?;
Ok(Some(payload))
}
#[cfg(test)]
mod tests {
use super::*;
use std::io::Cursor;
#[test]
fn round_trips_a_payload() {
let mut buf = Vec::new();
write_frame(&mut buf, b"hello").unwrap();
assert_eq!(buf, [0, 0, 0, 5, b'h', b'e', b'l', b'l', b'o']);
let read = read_frame(Cursor::new(buf)).unwrap();
assert_eq!(read.as_deref(), Some(&b"hello"[..]));
}
#[test]
fn clean_end_of_stream_is_none() {
assert_eq!(read_frame(Cursor::new(Vec::<u8>::new())).unwrap(), None);
}
#[test]
fn truncated_length_is_an_error() {
let err = read_frame(Cursor::new(vec![0, 0])).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
}
#[test]
fn truncated_payload_is_an_error() {
let err = read_frame(Cursor::new(vec![0, 0, 0, 9, b'x'])).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::UnexpectedEof);
}
#[test]
fn oversized_length_is_rejected_before_allocating() {
let err = read_frame(Cursor::new(vec![0xff, 0xff, 0xff, 0xff])).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
#[test]
fn header_encodes_the_length_as_big_endian() {
assert_eq!(header(5).unwrap(), [0, 0, 0, 5]);
}
#[test]
fn payload_len_rejects_a_frame_over_the_limit() {
let err = payload_len([0xff; 4]).unwrap_err();
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
}
}