use std::{
error, fmt,
io::{self, Read, Write},
};
use prost::Message;
use crate::decode_budget::with_decode_budget;
pub const MAX_FRAME_LEN: u32 = 256 * 1024 * 1024;
pub const DEFAULT_MAX_DECODE_BYTES: usize = 4 * MAX_FRAME_LEN as usize;
#[derive(Debug)]
pub enum FrameError {
Io(io::Error),
Decode(prost::DecodeError),
FrameTooLarge {
len: u32,
max: u32,
},
Truncated,
}
impl fmt::Display for FrameError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Io(e) => write!(f, "frame I/O error: {e}"),
Self::Decode(e) => write!(f, "frame decode error: {e}"),
Self::FrameTooLarge { len, max } => write!(f, "frame of {len} bytes exceeds maximum of {max} bytes"),
Self::Truncated => f.write_str("stream ended mid-frame"),
}
}
}
impl error::Error for FrameError {}
impl From<io::Error> for FrameError {
fn from(e: io::Error) -> Self {
Self::Io(e)
}
}
pub fn exceeds_max_frame_len(msg: &impl Message) -> Option<u32> {
let len = u32::try_from(msg.encoded_len()).unwrap_or(u32::MAX);
(len > MAX_FRAME_LEN).then_some(len)
}
pub fn write_frame(writer: &mut impl Write, msg: &impl Message) -> Result<(), FrameError> {
let body = encode_to_capped_vec(msg)?;
let len = u32::try_from(body.len()).unwrap_or(u32::MAX);
writer.write_all(&len.to_le_bytes())?;
writer.write_all(&body)?;
writer.flush()?;
Ok(())
}
pub fn encode_to_capped_vec(msg: &impl Message) -> Result<Vec<u8>, FrameError> {
let encoded_len = msg.encoded_len();
u32::try_from(encoded_len)
.ok()
.filter(|&len| len <= MAX_FRAME_LEN)
.ok_or(FrameError::FrameTooLarge {
len: u32::try_from(encoded_len).unwrap_or(u32::MAX),
max: MAX_FRAME_LEN,
})?;
Ok(msg.encode_to_vec())
}
pub fn encode_framed_into(msg: &impl Message, buf: &mut Vec<u8>) -> Result<(), FrameError> {
let encoded_len = msg.encoded_len();
let len = u32::try_from(encoded_len)
.ok()
.filter(|&len| len <= MAX_FRAME_LEN)
.ok_or(FrameError::FrameTooLarge {
len: u32::try_from(encoded_len).unwrap_or(u32::MAX),
max: MAX_FRAME_LEN,
})?;
buf.clear();
buf.reserve(4 + encoded_len);
buf.extend_from_slice(&len.to_le_bytes());
msg.encode_raw(buf);
Ok(())
}
pub fn decode_frame<M: Message + Default>(bytes: &[u8]) -> Result<M, FrameError> {
if bytes.len() > MAX_FRAME_LEN as usize {
return Err(FrameError::FrameTooLarge {
len: u32::try_from(bytes.len()).unwrap_or(u32::MAX),
max: MAX_FRAME_LEN,
});
}
with_decode_budget(DEFAULT_MAX_DECODE_BYTES, || M::decode(bytes)).map_err(FrameError::Decode)
}
#[derive(Debug)]
pub struct FrameReader<R: Read> {
inner: R,
max_frame_len: u32,
}
impl<R: Read> FrameReader<R> {
pub fn new(inner: R) -> Self {
Self {
inner,
max_frame_len: MAX_FRAME_LEN,
}
}
pub fn with_max_frame_len(inner: R, max_frame_len: u32) -> Self {
Self { inner, max_frame_len }
}
pub fn read<M: Message + Default>(&mut self) -> Result<Option<M>, FrameError> {
let mut len_bytes = [0u8; 4];
match read_exact_or_eof(&mut self.inner, &mut len_bytes)? {
ReadOutcome::CleanEof => return Ok(None),
ReadOutcome::Truncated => return Err(FrameError::Truncated),
ReadOutcome::Filled => {}
}
let len = u32::from_le_bytes(len_bytes);
if len > self.max_frame_len {
return Err(FrameError::FrameTooLarge {
len,
max: self.max_frame_len,
});
}
let mut body = vec![0u8; len as usize];
match read_exact_or_eof(&mut self.inner, &mut body)? {
ReadOutcome::Filled => {}
ReadOutcome::CleanEof | ReadOutcome::Truncated => return Err(FrameError::Truncated),
}
with_decode_budget(DEFAULT_MAX_DECODE_BYTES, || M::decode(body.as_slice()))
.map(Some)
.map_err(FrameError::Decode)
}
}
enum ReadOutcome {
Filled,
CleanEof,
Truncated,
}
fn read_exact_or_eof(reader: &mut impl Read, buf: &mut [u8]) -> io::Result<ReadOutcome> {
let mut filled = 0;
while filled < buf.len() {
match reader.read(&mut buf[filled..]) {
Ok(0) => {
return Ok(if filled == 0 {
ReadOutcome::CleanEof
} else {
ReadOutcome::Truncated
});
}
Ok(n) => filled += n,
Err(e) if e.kind() == io::ErrorKind::Interrupted => {}
Err(e) => return Err(e),
}
}
Ok(ReadOutcome::Filled)
}