use serde::Serialize;
use serde::de::DeserializeOwned;
use std::fmt;
use std::io::{Error, ErrorKind, Result};
use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt};
pub const MAX_FRAME_BYTES: usize = 8 * 1024 * 1024;
pub const MAX_PARTIAL_FRAME_BYTES: usize = 16 * 1024 * 1024;
const _: () = assert!(MAX_FRAME_BYTES + 4 <= MAX_PARTIAL_FRAME_BYTES);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct TooLarge {
pub len: u64,
pub max: usize,
}
impl fmt::Display for TooLarge {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"frame of {} bytes exceeds {max}",
self.len,
max = self.max
)
}
}
impl std::error::Error for TooLarge {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UnexpectedEof {
pub expected: u64,
pub got: u64,
}
impl fmt::Display for UnexpectedEof {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"stream ended mid-frame: {} of {} bytes arrived",
self.got, self.expected
)
}
}
impl std::error::Error for UnexpectedEof {}
fn too_large(len: u64, max: usize) -> Error {
Error::new(ErrorKind::InvalidInput, TooLarge { len, max })
}
pub async fn write_frame<W, T>(w: &mut W, value: &T) -> Result<()>
where
W: AsyncWrite + Unpin,
T: Serialize + ?Sized,
{
let body = serde_json::to_vec(value).map_err(Error::other)?;
let len = body.len();
if len > MAX_FRAME_BYTES {
return Err(too_large(len as u64, MAX_FRAME_BYTES));
}
let prefix = u32::try_from(len).map_err(|_| too_large(len as u64, MAX_FRAME_BYTES))?;
let mut buf = Vec::with_capacity(4 + len);
buf.extend_from_slice(&prefix.to_be_bytes());
buf.extend_from_slice(&body);
w.write_all(&buf).await?;
w.flush().await
}
pub async fn read_frame<R, T>(r: &mut R) -> Result<Option<T>>
where
R: AsyncRead + Unpin,
T: DeserializeOwned,
{
let mut header = [0u8; 4];
if !fill(r, &mut header, 4).await? {
return Ok(None);
}
let len = u32::from_be_bytes(header) as u64;
if len > MAX_FRAME_BYTES as u64 {
return Err(too_large(len, MAX_FRAME_BYTES));
}
let mut body = vec![0u8; len as usize];
if !body.is_empty() && !fill(r, &mut body, len).await? {
return Err(Error::new(
ErrorKind::UnexpectedEof,
UnexpectedEof {
expected: len,
got: 0,
},
));
}
serde_json::from_slice::<T>(&body)
.map(Some)
.map_err(|e| Error::new(ErrorKind::InvalidData, format!("bad frame json: {e}")))
}
async fn fill<R>(r: &mut R, dst: &mut [u8], announced: u64) -> Result<bool>
where
R: AsyncRead + Unpin,
{
if dst.len() > MAX_PARTIAL_FRAME_BYTES {
return Err(too_large(dst.len() as u64, MAX_PARTIAL_FRAME_BYTES));
}
let mut got = 0usize;
while got < dst.len() {
match r.read(&mut dst[got..]).await {
Ok(0) => {
if got == 0 {
return Ok(false);
}
return Err(Error::new(
ErrorKind::UnexpectedEof,
UnexpectedEof {
expected: announced,
got: got as u64,
},
));
}
Ok(n) => got += n,
Err(e) if e.kind() == ErrorKind::Interrupted => continue,
Err(e) => return Err(e),
}
}
Ok(true)
}
const READ_CHUNK: usize = 64 * 1024;
#[derive(Debug, Default)]
pub struct FrameReader {
buf: Vec<u8>,
}
impl FrameReader {
pub fn new() -> Self {
Self::default()
}
pub async fn next<R, T>(&mut self, r: &mut R) -> Result<Option<T>>
where
R: AsyncRead + Unpin,
T: DeserializeOwned,
{
loop {
if let Some(end) = self.frame_end()? {
let decoded = serde_json::from_slice::<T>(&self.buf[4..end]);
self.buf.drain(..end);
return decoded.map(Some).map_err(|e| {
Error::new(ErrorKind::InvalidData, format!("bad frame json: {e}"))
});
}
self.buf.reserve(READ_CHUNK);
if r.read_buf(&mut self.buf).await? == 0 {
if self.buf.is_empty() {
return Ok(None);
}
let got = self.buf.len() as u64;
let eof = match self.announced() {
Some(len) => UnexpectedEof {
expected: len,
got: got - 4,
},
None => UnexpectedEof { expected: 4, got },
};
return Err(Error::new(ErrorKind::UnexpectedEof, eof));
}
}
}
fn announced(&self) -> Option<u64> {
let header: [u8; 4] = self.buf.get(..4)?.try_into().ok()?;
Some(u32::from_be_bytes(header) as u64)
}
fn frame_end(&self) -> Result<Option<usize>> {
let Some(len) = self.announced() else {
return Ok(None);
};
if len > MAX_FRAME_BYTES as u64 {
return Err(too_large(len, MAX_FRAME_BYTES));
}
let end = 4 + len as usize;
Ok((self.buf.len() >= end).then_some(end))
}
}
pub fn is_too_large(err: &Error) -> bool {
err.get_ref()
.and_then(|r| r.downcast_ref::<TooLarge>())
.is_some()
}
pub fn is_bad_frame(err: &Error) -> bool {
err.kind() == ErrorKind::InvalidData
}