use std::io::{self, Read, Write};
use crate::Error;
pub struct Crc32SegmentReader<R: Read> {
reader: R,
buffer: Vec<u8>,
}
impl<R: Read> Crc32SegmentReader<R> {
pub fn new(reader: R) -> Self {
Self {
reader,
buffer: Vec::new(),
}
}
pub fn next_record(&mut self) -> Result<Option<&[u8]>, Error> {
let mut len_buf = [0u8; 4];
match self.reader.read(&mut len_buf)? {
0 => return Ok(None),
n if n < 4 => return Err(Error::TruncatedFrame),
_ => {}
}
let len = u32::from_be_bytes(len_buf) as usize;
self.buffer.resize(len, 0);
self.reader.read_exact(&mut self.buffer).map_err(map_eof)?;
let mut crc_buf = [0u8; 4];
self.reader.read_exact(&mut crc_buf).map_err(map_eof)?;
let expected = u32::from_be_bytes(crc_buf);
let actual = crc32c::crc32c(&self.buffer);
if expected != actual {
return Err(Error::ChecksumMismatch);
}
Ok(Some(&self.buffer))
}
}
pub struct Crc32SegmentWriter<W: Write> {
writer: W,
}
impl<W: Write> Crc32SegmentWriter<W> {
pub fn new(writer: W) -> Self {
Self { writer }
}
pub fn write(&mut self, record: &[u8]) -> io::Result<()> {
let len = record.len() as u32;
self.writer.write_all(&len.to_be_bytes())?;
self.writer.write_all(record)?;
let crc = crc32c::crc32c(record);
self.writer.write_all(&crc.to_be_bytes())?;
Ok(())
}
}
fn map_eof(e: io::Error) -> Error {
if e.kind() == io::ErrorKind::UnexpectedEof {
Error::TruncatedFrame
} else {
Error::Io(e)
}
}
#[cfg(test)]
mod tests {
use super::*;
fn build(records: &[&[u8]]) -> Vec<u8> {
let mut buf = Vec::new();
let mut w = Crc32SegmentWriter::new(&mut buf);
for r in records {
w.write(r).unwrap();
}
buf
}
#[test]
fn round_trip_with_checksum() {
let buf = build(&[b"alice", b"bob", b"carol"]);
let mut r = Crc32SegmentReader::new(buf.as_slice());
assert_eq!(r.next_record().unwrap().unwrap(), b"alice");
assert_eq!(r.next_record().unwrap().unwrap(), b"bob");
assert_eq!(r.next_record().unwrap().unwrap(), b"carol");
assert!(r.next_record().unwrap().is_none());
}
#[test]
fn empty_segment_yields_none() {
let mut r = Crc32SegmentReader::new(&[][..]);
assert!(r.next_record().unwrap().is_none());
}
#[test]
fn corrupted_payload_detected() {
let mut buf = build(&[b"hello"]);
buf[4] ^= 0x80;
let mut r = Crc32SegmentReader::new(buf.as_slice());
assert!(matches!(r.next_record(), Err(Error::ChecksumMismatch)));
}
#[test]
fn corrupted_trailer_detected() {
let mut buf = build(&[b"hello"]);
let last = buf.len() - 1;
buf[last] ^= 0xff;
let mut r = Crc32SegmentReader::new(buf.as_slice());
assert!(matches!(r.next_record(), Err(Error::ChecksumMismatch)));
}
#[test]
fn truncated_trailer_surfaces_typed_error() {
let mut buf = build(&[b"hello"]);
buf.truncate(buf.len() - 2); let mut r = Crc32SegmentReader::new(buf.as_slice());
assert!(matches!(r.next_record(), Err(Error::TruncatedFrame)));
}
#[test]
fn truncated_payload_surfaces_typed_error() {
let mut buf = Vec::new();
buf.extend_from_slice(&[0, 0, 0, 10]); buf.extend_from_slice(b"abc");
let mut r = Crc32SegmentReader::new(buf.as_slice());
assert!(matches!(r.next_record(), Err(Error::TruncatedFrame)));
}
#[test]
fn zero_length_record_round_trips() {
let buf = build(&[&[], b"after-empty"]);
let mut r = Crc32SegmentReader::new(buf.as_slice());
assert_eq!(r.next_record().unwrap().unwrap(), b"");
assert_eq!(r.next_record().unwrap().unwrap(), b"after-empty");
}
}