use crate::core::{Error, Result};
pub trait Frame: Send + Sync + 'static {
fn next_record_len(&mut self, buf: &[u8]) -> Result<Option<usize>>;
}
pub struct FixedLength {
record_len: usize,
}
impl FixedLength {
pub fn new(record_len: usize) -> Self {
assert!(record_len > 0, "record_len must be > 0");
Self { record_len }
}
}
impl Frame for FixedLength {
fn next_record_len(&mut self, buf: &[u8]) -> Result<Option<usize>> {
if buf.len() >= self.record_len {
Ok(Some(self.record_len))
} else {
Ok(None)
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum PrefixWidth {
U8,
U16Be,
U32Be,
}
impl PrefixWidth {
fn byte_count(self) -> usize {
match self {
Self::U8 => 1,
Self::U16Be => 2,
Self::U32Be => 4,
}
}
fn read(self, buf: &[u8]) -> u64 {
match self {
Self::U8 => buf[0] as u64,
Self::U16Be => u16::from_be_bytes([buf[0], buf[1]]) as u64,
Self::U32Be => u32::from_be_bytes([buf[0], buf[1], buf[2], buf[3]]) as u64,
}
}
}
pub struct LengthPrefixed {
width: PrefixWidth,
max_payload: usize,
}
impl LengthPrefixed {
pub fn new(width: PrefixWidth, max_payload: usize) -> Self {
Self { width, max_payload }
}
}
impl Frame for LengthPrefixed {
fn next_record_len(&mut self, buf: &[u8]) -> Result<Option<usize>> {
let header = self.width.byte_count();
if buf.len() < header {
return Ok(None);
}
let payload_len = self.width.read(&buf[..header]) as usize;
if payload_len > self.max_payload {
return Err(Error::new(
crate::core::ErrorKind::Decode,
format!(
"framing: payload length {payload_len} exceeds max {}",
self.max_payload
),
));
}
let total = header + payload_len;
if buf.len() >= total {
Ok(Some(total))
} else {
Ok(None)
}
}
}
pub struct Delimiter {
byte: u8,
max_len: usize,
}
impl Delimiter {
pub fn new(byte: u8, max_len: usize) -> Self {
assert!(max_len > 0, "max_len must be > 0");
Self { byte, max_len }
}
}
impl Frame for Delimiter {
fn next_record_len(&mut self, buf: &[u8]) -> Result<Option<usize>> {
let scan_len = buf.len().min(self.max_len);
match buf[..scan_len].iter().position(|&b| b == self.byte) {
Some(pos) => Ok(Some(pos + 1)), None if buf.len() >= self.max_len => Err(Error::new(
crate::core::ErrorKind::Decode,
format!("framing: no delimiter found within {} bytes", self.max_len),
)),
None => Ok(None), }
}
}
pub struct Custom<F>
where
F: FnMut(&[u8]) -> Result<Option<usize>> + Send + Sync + 'static,
{
f: F,
}
impl<F> Custom<F>
where
F: FnMut(&[u8]) -> Result<Option<usize>> + Send + Sync + 'static,
{
pub fn new(f: F) -> Self {
Self { f }
}
}
impl<F> Frame for Custom<F>
where
F: FnMut(&[u8]) -> Result<Option<usize>> + Send + Sync + 'static,
{
fn next_record_len(&mut self, buf: &[u8]) -> Result<Option<usize>> {
(self.f)(buf)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fixed_needs_full_record() {
let mut f = FixedLength::new(8);
assert_eq!(f.next_record_len(b"hello").unwrap(), None);
assert_eq!(f.next_record_len(b"hello!!!").unwrap(), Some(8));
assert_eq!(f.next_record_len(b"hello!!!!extra").unwrap(), Some(8));
}
#[test]
fn length_prefixed_u8() {
let mut f = LengthPrefixed::new(PrefixWidth::U8, 256);
let buf: &[u8] = &[3, b'a', b'b', b'c'];
assert_eq!(f.next_record_len(buf).unwrap(), Some(4));
}
#[test]
fn length_prefixed_too_short() {
let mut f = LengthPrefixed::new(PrefixWidth::U16Be, 1024);
assert_eq!(f.next_record_len(&[0x00]).unwrap(), None);
}
#[test]
fn length_prefixed_exceeds_max() {
let mut f = LengthPrefixed::new(PrefixWidth::U8, 4);
let buf: &[u8] = &[10, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0]; assert!(f.next_record_len(buf).is_err());
}
#[test]
fn delimiter_finds_newline() {
let mut f = Delimiter::new(b'\n', 1024);
assert_eq!(f.next_record_len(b"hello\nworld").unwrap(), Some(6));
}
#[test]
fn delimiter_needs_more_data() {
let mut f = Delimiter::new(b'\n', 1024);
assert_eq!(f.next_record_len(b"no newline here").unwrap(), None);
}
#[test]
fn delimiter_exceeds_max() {
let mut f = Delimiter::new(b'\n', 5);
assert!(f.next_record_len(b"abcdef").is_err());
}
#[test]
fn custom_framer() {
let mut f = Custom::new(|buf: &[u8]| {
if buf.len() >= 3 {
Ok(Some(3))
} else {
Ok(None)
}
});
assert_eq!(f.next_record_len(b"ab").unwrap(), None);
assert_eq!(f.next_record_len(b"abc").unwrap(), Some(3));
}
}