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;
const MAX_RETAINED_BODY: usize = 64 * 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 frames = FrameReader::new();
loop {
match frames.poll::<M, R>(reader, max_frame_len)? {
FramePoll::Frame(msg) => return Ok(Some(msg)),
FramePoll::Eof => return Ok(None),
FramePoll::Progress => {}
FramePoll::WouldBlock { .. } => {
return Err(FrameError::Io(io::Error::new(
io::ErrorKind::WouldBlock,
"read timed out before a complete frame",
)));
}
}
}
}
#[derive(Debug)]
pub enum FramePoll<M> {
Frame(M),
Eof,
Progress,
WouldBlock { in_progress: bool },
}
#[derive(Default)]
pub struct FrameReader {
len_buf: [u8; 4],
len_filled: usize,
body: Vec<u8>,
body_filled: usize,
}
enum ReadStep {
Progress,
WouldBlock,
Eof,
}
impl FrameReader {
pub fn new() -> FrameReader {
FrameReader::default()
}
pub fn poll<M: Message, R: Read>(
&mut self,
reader: &mut R,
max_frame_len: u32,
) -> Result<FramePoll<M>, FrameError> {
if self.len_filled < 4 {
match read_once(reader, &mut self.len_buf, &mut self.len_filled)? {
ReadStep::Eof => {
return if self.len_filled == 0 {
Ok(FramePoll::Eof)
} else {
Err(torn("eof partway through a frame length prefix"))
};
}
ReadStep::WouldBlock => {
return Ok(FramePoll::WouldBlock {
in_progress: self.len_filled > 0,
});
}
ReadStep::Progress => {
if self.len_filled < 4 {
return Ok(FramePoll::Progress);
}
let len = u32::from_be_bytes(self.len_buf);
if len > max_frame_len {
return Err(FrameError::TooLarge {
len,
max: max_frame_len,
});
}
self.body.clear();
self.body.resize(len as usize, 0);
self.body_filled = 0;
if len == 0 {
return self.finish::<M>();
}
return Ok(FramePoll::Progress);
}
}
}
match read_once(reader, &mut self.body, &mut self.body_filled)? {
ReadStep::Eof => Err(torn("eof partway through a frame body")),
ReadStep::WouldBlock => Ok(FramePoll::WouldBlock { in_progress: true }),
ReadStep::Progress => {
if self.body_filled < self.body.len() {
Ok(FramePoll::Progress)
} else {
self.finish::<M>()
}
}
}
}
fn finish<M: Message>(&mut self) -> Result<FramePoll<M>, FrameError> {
let msg = M::parse(&self.body).map_err(FrameError::Parse)?;
self.len_filled = 0;
self.body.clear();
if self.body.capacity() > MAX_RETAINED_BODY {
self.body.shrink_to(MAX_RETAINED_BODY);
}
self.body_filled = 0;
Ok(FramePoll::Frame(msg))
}
}
fn torn(msg: &'static str) -> FrameError {
FrameError::Io(io::Error::new(io::ErrorKind::UnexpectedEof, msg))
}
fn read_once<R: Read>(
reader: &mut R,
buf: &mut [u8],
filled: &mut usize,
) -> Result<ReadStep, FrameError> {
loop {
match reader.read(&mut buf[*filled..]) {
Ok(0) => return Ok(ReadStep::Eof),
Ok(n) => {
*filled += n;
return Ok(ReadStep::Progress);
}
Err(ref err) if err.kind() == io::ErrorKind::Interrupted => {}
Err(ref err)
if err.kind() == io::ErrorKind::WouldBlock
|| err.kind() == io::ErrorKind::TimedOut =>
{
return Ok(ReadStep::WouldBlock);
}
Err(err) => return Err(FrameError::Io(err)),
}
}
}
#[cfg(feature = "tokio")]
pub async fn write_frame_async<M, W>(
writer: &mut W,
msg: &M,
max_frame_len: u32,
) -> Result<(), FrameError>
where
M: Message,
W: tokio::io::AsyncWrite + Unpin,
{
use tokio::io::AsyncWriteExt;
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()).await?;
writer.write_all(&body).await?;
Ok(())
}
#[cfg(feature = "tokio")]
pub async fn read_frame_async<M, R>(
reader: &mut R,
max_frame_len: u32,
) -> Result<Option<M>, FrameError>
where
M: Message,
R: tokio::io::AsyncRead + Unpin,
{
let mut len_buf = [0u8; 4];
if !read_full_or_eof_async(reader, &mut len_buf).await? {
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];
read_exact_async(reader, &mut body).await?;
let msg = M::parse(&body).map_err(FrameError::Parse)?;
Ok(Some(msg))
}
#[cfg(feature = "tokio")]
async fn read_full_or_eof_async<R>(reader: &mut R, buf: &mut [u8]) -> io::Result<bool>
where
R: tokio::io::AsyncRead + Unpin,
{
use tokio::io::AsyncReadExt;
let mut filled = 0;
while filled < buf.len() {
match reader.read(&mut buf[filled..]).await {
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(err) => return Err(err),
}
}
Ok(true)
}
#[cfg(feature = "tokio")]
async fn read_exact_async<R>(reader: &mut R, buf: &mut [u8]) -> io::Result<()>
where
R: tokio::io::AsyncRead + Unpin,
{
use tokio::io::AsyncReadExt;
let mut filled = 0;
while filled < buf.len() {
match reader.read(&mut buf[filled..]).await {
Ok(0) => {
return Err(io::Error::new(
io::ErrorKind::UnexpectedEof,
"eof partway through a frame body",
));
}
Ok(n) => filled += n,
Err(err) => return Err(err),
}
}
Ok(())
}
#[cfg(test)]
mod frame_reader_tests {
use std::collections::VecDeque;
use protobuf::Serialize;
use super::*;
use crate::tephra::Request;
struct MockReader {
events: VecDeque<Event>,
buf: Vec<u8>,
pos: usize,
}
enum Event {
Data(Vec<u8>),
Block,
}
impl MockReader {
fn new(events: Vec<Event>) -> MockReader {
MockReader {
events: events.into(),
buf: Vec::new(),
pos: 0,
}
}
}
impl Read for MockReader {
fn read(&mut self, out: &mut [u8]) -> io::Result<usize> {
if self.pos < self.buf.len() {
let n = (self.buf.len() - self.pos).min(out.len());
out[..n].copy_from_slice(&self.buf[self.pos..self.pos + n]);
self.pos += n;
return Ok(n);
}
match self.events.pop_front() {
Some(Event::Data(data)) => {
self.buf = data;
self.pos = 0;
self.read(out)
}
Some(Event::Block) => Err(io::Error::from(io::ErrorKind::WouldBlock)),
None => Ok(0),
}
}
}
fn framed(request_id: u64) -> Vec<u8> {
let mut request = Request::new();
request.set_request_id(request_id);
let body = request.serialize().unwrap();
let mut out = (body.len() as u32).to_be_bytes().to_vec();
out.extend_from_slice(&body);
out
}
fn framed_large(request_id: u64, payload_len: usize) -> Vec<u8> {
let mut event = crate::tephra::Event::new();
event.set_payload(vec![0u8; payload_len]);
let mut append = crate::tephra::AppendRequest::new();
append.events_mut().push(event);
let mut request = Request::new();
request.set_request_id(request_id);
request.set_append(append);
let body = request.serialize().unwrap();
let mut out = (body.len() as u32).to_be_bytes().to_vec();
out.extend_from_slice(&body);
out
}
fn drive_to_frame(reader: &mut FrameReader, mock: &mut MockReader) -> u64 {
loop {
match reader
.poll::<Request, _>(mock, DEFAULT_MAX_FRAME_LEN)
.unwrap()
{
FramePoll::Frame(req) => return req.request_id(),
FramePoll::Progress | FramePoll::WouldBlock { .. } => {}
FramePoll::Eof => panic!("unexpected eof before a frame"),
}
}
}
#[test]
fn reads_a_whole_frame() {
let mut mock = MockReader::new(vec![Event::Data(framed(7))]);
let mut reader = FrameReader::new();
assert_eq!(drive_to_frame(&mut reader, &mut mock), 7);
}
#[test]
fn poll_yields_after_each_read_so_a_trickle_is_observable() {
let frame = framed(21);
let total = frame.len();
let events = frame.iter().map(|b| Event::Data(vec![*b])).collect();
let mut mock = MockReader::new(events);
let mut reader = FrameReader::new();
let mut progress = 0;
loop {
match reader
.poll::<Request, _>(&mut mock, DEFAULT_MAX_FRAME_LEN)
.unwrap()
{
FramePoll::Frame(req) => {
assert_eq!(req.request_id(), 21);
break;
}
FramePoll::Progress => progress += 1,
other => panic!("expected progress, got {other:?}"),
}
}
assert_eq!(
progress,
total - 1,
"expected one poll per byte for {total} bytes"
);
}
#[test]
fn would_block_is_idle_before_the_first_byte() {
let mut mock = MockReader::new(vec![Event::Block, Event::Data(framed(9))]);
let mut reader = FrameReader::new();
match reader
.poll::<Request, _>(&mut mock, DEFAULT_MAX_FRAME_LEN)
.unwrap()
{
FramePoll::WouldBlock { in_progress } => assert!(!in_progress, "idle at a boundary"),
other => panic!("expected a boundary would-block, got {other:?}"),
}
assert_eq!(drive_to_frame(&mut reader, &mut mock), 9);
}
#[test]
fn resumes_across_a_would_block_mid_length_prefix() {
let frame = framed(11);
let (head, tail) = frame.split_at(2);
let mut mock = MockReader::new(vec![
Event::Data(head.to_vec()),
Event::Block,
Event::Data(tail.to_vec()),
]);
let mut reader = FrameReader::new();
assert!(matches!(
reader.poll::<Request, _>(&mut mock, DEFAULT_MAX_FRAME_LEN),
Ok(FramePoll::Progress)
));
match reader
.poll::<Request, _>(&mut mock, DEFAULT_MAX_FRAME_LEN)
.unwrap()
{
FramePoll::WouldBlock { in_progress } => assert!(in_progress, "partial length prefix"),
other => panic!("expected an in-progress would-block, got {other:?}"),
}
assert_eq!(drive_to_frame(&mut reader, &mut mock), 11);
}
#[test]
fn resumes_across_a_would_block_mid_body() {
let frame = framed(13);
let (head, tail) = frame.split_at(5);
let mut mock = MockReader::new(vec![
Event::Data(head.to_vec()),
Event::Block,
Event::Data(tail.to_vec()),
]);
let mut reader = FrameReader::new();
for _ in 0..2 {
assert!(matches!(
reader.poll::<Request, _>(&mut mock, DEFAULT_MAX_FRAME_LEN),
Ok(FramePoll::Progress)
));
}
match reader
.poll::<Request, _>(&mut mock, DEFAULT_MAX_FRAME_LEN)
.unwrap()
{
FramePoll::WouldBlock { in_progress } => assert!(in_progress, "partial body"),
other => panic!("expected an in-progress would-block, got {other:?}"),
}
assert_eq!(drive_to_frame(&mut reader, &mut mock), 13);
}
#[test]
fn shrinks_the_body_buffer_after_an_oversized_frame() {
let large = 512 * 1024;
let mut mock = MockReader::new(vec![Event::Data(framed_large(1, large))]);
let mut reader = FrameReader::new();
assert_eq!(drive_to_frame(&mut reader, &mut mock), 1);
assert!(
reader.body.capacity() < large,
"body buffer retained {} bytes after a {large}-byte frame",
reader.body.capacity(),
);
}
#[test]
fn reuses_its_buffer_across_frames() {
let mut bytes = framed(1);
bytes.extend_from_slice(&framed(2));
let mut mock = MockReader::new(vec![Event::Data(bytes)]);
let mut reader = FrameReader::new();
assert_eq!(drive_to_frame(&mut reader, &mut mock), 1);
assert_eq!(drive_to_frame(&mut reader, &mut mock), 2);
}
#[test]
fn clean_eof_at_a_boundary() {
let mut mock = MockReader::new(vec![]);
let mut reader = FrameReader::new();
assert!(matches!(
reader.poll::<Request, _>(&mut mock, DEFAULT_MAX_FRAME_LEN),
Ok(FramePoll::Eof)
));
}
#[test]
fn a_torn_frame_is_an_error() {
let frame = framed(15);
let len_prefix = frame[..4].to_vec();
let mut mock = MockReader::new(vec![Event::Data(len_prefix)]);
let mut reader = FrameReader::new();
let mut err = None;
for _ in 0..3 {
match reader.poll::<Request, _>(&mut mock, DEFAULT_MAX_FRAME_LEN) {
Ok(FramePoll::Progress | FramePoll::WouldBlock { .. }) => {}
Ok(other) => panic!("expected an error, got {other:?}"),
Err(e) => {
err = Some(e);
break;
}
}
}
assert!(
matches!(err, Some(FrameError::Io(_))),
"torn frame is an i/o error, got {err:?}"
);
}
#[test]
fn an_oversized_length_is_rejected() {
let mut prefix = 100u32.to_be_bytes().to_vec();
prefix.extend_from_slice(&[0u8; 4]);
let mut mock = MockReader::new(vec![Event::Data(prefix)]);
let mut reader = FrameReader::new();
let err = reader.poll::<Request, _>(&mut mock, 16).unwrap_err();
assert!(matches!(err, FrameError::TooLarge { len: 100, max: 16 }));
}
}