use bytes::{Buf, Bytes, BytesMut};
use crate::{Error, Result};
pub const DEFAULT_MAX_FRAME: usize = 100 * 1024 * 1024;
#[derive(Debug)]
pub struct FrameDecoder {
buf: BytesMut,
max_frame: usize,
}
impl Default for FrameDecoder {
fn default() -> Self {
Self::new(DEFAULT_MAX_FRAME)
}
}
impl FrameDecoder {
#[must_use]
pub fn new(max_frame: usize) -> Self {
Self {
buf: BytesMut::new(),
max_frame,
}
}
pub fn push(&mut self, bytes: &[u8]) {
self.buf.extend_from_slice(bytes);
}
#[must_use]
pub fn buffered(&self) -> usize {
self.buf.len()
}
#[must_use]
pub fn needed(&self) -> usize {
if self.buf.len() < 4 {
return 4 - self.buf.len();
}
let len = i32::from_be_bytes([self.buf[0], self.buf[1], self.buf[2], self.buf[3]]);
let len = usize::try_from(len).unwrap_or(0).min(self.max_frame);
(4 + len).saturating_sub(self.buf.len())
}
pub fn next_frame(&mut self) -> Result<Option<Bytes>> {
if self.buf.len() < 4 {
return Ok(None);
}
let len = i32::from_be_bytes([self.buf[0], self.buf[1], self.buf[2], self.buf[3]]);
let len = usize::try_from(len).map_err(|_| Error::FrameTooLarge {
len: usize::MAX,
limit: self.max_frame,
})?;
if len > self.max_frame {
return Err(Error::FrameTooLarge {
len,
limit: self.max_frame,
});
}
if self.buf.len() < 4 + len {
return Ok(None);
}
self.buf.advance(4);
Ok(Some(self.buf.split_to(len).freeze()))
}
}
pub fn frame(body: &[u8]) -> Result<Bytes> {
let len = i32::try_from(body.len()).map_err(|_| Error::FrameTooLarge {
len: body.len(),
limit: i32::MAX as usize,
})?;
let mut out = BytesMut::with_capacity(body.len() + 4);
out.extend_from_slice(&len.to_be_bytes());
out.extend_from_slice(body);
Ok(out.freeze())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_whole_frame_round_trips() {
let mut dec = FrameDecoder::default();
dec.push(&frame(b"hello").unwrap());
assert_eq!(dec.next_frame().unwrap().as_deref(), Some(&b"hello"[..]));
assert!(dec.next_frame().unwrap().is_none());
}
#[test]
fn a_frame_split_byte_by_byte_still_arrives() {
let bytes = frame(b"a longer payload").unwrap();
let mut dec = FrameDecoder::default();
for (i, b) in bytes.iter().enumerate() {
dec.push(&[*b]);
if i + 1 < bytes.len() {
assert!(
dec.next_frame().unwrap().is_none(),
"yielded a frame after {} of {} bytes",
i + 1,
bytes.len()
);
}
}
assert_eq!(
dec.next_frame().unwrap().as_deref(),
Some(&b"a longer payload"[..])
);
}
#[test]
fn many_frames_and_a_fragment_in_one_read() {
let mut wire = Vec::new();
wire.extend_from_slice(&frame(b"one").unwrap());
wire.extend_from_slice(&frame(b"two").unwrap());
let third = frame(b"three").unwrap();
wire.extend_from_slice(&third[..4]);
let mut dec = FrameDecoder::default();
dec.push(&wire);
assert_eq!(dec.next_frame().unwrap().as_deref(), Some(&b"one"[..]));
assert_eq!(dec.next_frame().unwrap().as_deref(), Some(&b"two"[..]));
assert!(dec.next_frame().unwrap().is_none());
dec.push(&third[4..]);
assert_eq!(dec.next_frame().unwrap().as_deref(), Some(&b"three"[..]));
}
#[test]
fn an_empty_frame_is_a_frame() {
let mut dec = FrameDecoder::default();
dec.push(&frame(b"").unwrap());
assert_eq!(dec.next_frame().unwrap().as_deref(), Some(&b""[..]));
}
#[test]
fn an_absurd_length_prefix_is_rejected_before_allocating() {
let mut dec = FrameDecoder::new(1024);
dec.push(&i32::MAX.to_be_bytes());
assert!(matches!(
dec.next_frame(),
Err(Error::FrameTooLarge { limit: 1024, .. })
));
assert_eq!(
dec.buffered(),
4,
"a rejected frame must not consume the buffer; the caller drops the connection"
);
}
#[test]
fn a_negative_length_prefix_is_rejected() {
let mut dec = FrameDecoder::default();
dec.push(&(-1i32).to_be_bytes());
assert!(matches!(dec.next_frame(), Err(Error::FrameTooLarge { .. })));
}
}