use std::io::{self, Read, Write};
use xxhash_rust::xxh3::xxh3_64;
use crate::Error;
pub struct Xxh3SegmentReader<R: Read> {
reader: R,
buffer: Vec<u8>,
}
impl<R: Read> Xxh3SegmentReader<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 hash_buf = [0u8; 8];
self.reader.read_exact(&mut hash_buf).map_err(map_eof)?;
let expected = u64::from_be_bytes(hash_buf);
let actual = xxh3_64(&self.buffer);
if expected != actual {
return Err(Error::ChecksumMismatch);
}
Ok(Some(&self.buffer))
}
}
pub struct Xxh3SegmentWriter<W: Write> {
writer: W,
}
impl<W: Write> Xxh3SegmentWriter<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 hash = xxh3_64(record);
self.writer.write_all(&hash.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 = Xxh3SegmentWriter::new(&mut buf);
for r in records {
w.write(r).unwrap();
}
buf
}
#[test]
fn round_trip_with_hash() {
let buf = build(&[b"alpha", b"beta", b"gamma"]);
let mut r = Xxh3SegmentReader::new(buf.as_slice());
assert_eq!(r.next_record().unwrap().unwrap(), b"alpha");
assert_eq!(r.next_record().unwrap().unwrap(), b"beta");
assert_eq!(r.next_record().unwrap().unwrap(), b"gamma");
assert!(r.next_record().unwrap().is_none());
}
#[test]
fn empty_segment_yields_none() {
let mut r = Xxh3SegmentReader::new(&[][..]);
assert!(r.next_record().unwrap().is_none());
}
#[test]
fn corrupted_payload_detected() {
let mut buf = build(&[b"hello"]);
buf[4] ^= 0x40;
let mut r = Xxh3SegmentReader::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 = Xxh3SegmentReader::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() - 4); let mut r = Xxh3SegmentReader::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 = Xxh3SegmentReader::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 = Xxh3SegmentReader::new(buf.as_slice());
assert_eq!(r.next_record().unwrap().unwrap(), b"");
assert_eq!(r.next_record().unwrap().unwrap(), b"after-empty");
}
}