use std::io::{self, Read, Write};
use std::{error, fmt};
use protobuf::{Message, ParseError, SerializeError};
pub const DEFAULT_MAX_FRAME_LEN: u32 = 16 * 1024 * 1024;
#[derive(Debug)]
pub enum FrameError {
Io(io::Error),
Parse(ParseError),
Serialize(SerializeError),
TooLarge { len: u32, max: u32 },
}
impl fmt::Display for FrameError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
FrameError::Io(err) => write!(f, "frame i/o error: {err}"),
FrameError::Parse(err) => write!(f, "frame decode error: {err}"),
FrameError::Serialize(err) => write!(f, "frame encode error: {err}"),
FrameError::TooLarge { len, max } => {
write!(f, "frame length {len} exceeds the maximum of {max}")
}
}
}
}
impl error::Error for FrameError {
fn source(&self) -> Option<&(dyn error::Error + 'static)> {
match self {
FrameError::Io(err) => Some(err),
_ => None,
}
}
}
impl From<io::Error> for FrameError {
fn from(err: io::Error) -> Self {
FrameError::Io(err)
}
}
pub fn write_frame<M: Message, W: Write>(
writer: &mut W,
msg: &M,
max_frame_len: u32,
) -> Result<(), FrameError> {
let body = msg.serialize().map_err(FrameError::Serialize)?;
if body.len() as u64 > u64::from(max_frame_len) {
return Err(FrameError::TooLarge {
len: body.len().min(u32::MAX as usize) as u32,
max: max_frame_len,
});
}
writer.write_all(&(body.len() as u32).to_be_bytes())?;
writer.write_all(&body)?;
Ok(())
}
pub fn read_frame<M: Message, R: Read>(
reader: &mut R,
max_frame_len: u32,
) -> Result<Option<M>, FrameError> {
let mut len_buf = [0u8; 4];
if !read_full_or_eof(reader, &mut len_buf)? {
return Ok(None);
}
let len = u32::from_be_bytes(len_buf);
if len > max_frame_len {
return Err(FrameError::TooLarge {
len,
max: max_frame_len,
});
}
let mut body = vec![0u8; len as usize];
reader.read_exact(&mut body)?;
let msg = M::parse(&body).map_err(FrameError::Parse)?;
Ok(Some(msg))
}
fn read_full_or_eof<R: Read>(reader: &mut R, buf: &mut [u8]) -> io::Result<bool> {
let mut filled = 0;
while filled < buf.len() {
match reader.read(&mut buf[filled..]) {
Ok(0) => {
if filled == 0 {
return Ok(false);
}
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"eof partway through a frame length prefix",
));
}
Ok(n) => filled += n,
Err(ref err) if err.kind() == io::ErrorKind::Interrupted => {}
Err(err) => return Err(err),
}
}
Ok(true)
}