#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(in crate::db) struct ByteDecodeError;
pub(in crate::db) struct ByteReader<'a> {
bytes: &'a [u8],
offset: usize,
}
impl<'a> ByteReader<'a> {
pub(in crate::db) const fn new(bytes: &'a [u8]) -> Self {
Self { bytes, offset: 0 }
}
pub(in crate::db) const fn remaining(&self) -> usize {
self.bytes.len().saturating_sub(self.offset)
}
pub(in crate::db) fn read_exact(&mut self, len: usize) -> Result<&'a [u8], ByteDecodeError> {
let end = self.offset.checked_add(len).ok_or(ByteDecodeError)?;
let value = self.bytes.get(self.offset..end).ok_or(ByteDecodeError)?;
self.offset = end;
Ok(value)
}
pub(in crate::db) fn read_array<const N: usize>(&mut self) -> Result<[u8; N], ByteDecodeError> {
self.read_exact(N)?.try_into().map_err(|_| ByteDecodeError)
}
pub(in crate::db) fn read_u8(&mut self) -> Result<u8, ByteDecodeError> {
Ok(self.read_array::<1>()?[0])
}
pub(in crate::db) fn read_u16(&mut self) -> Result<u16, ByteDecodeError> {
Ok(u16::from_be_bytes(self.read_array()?))
}
pub(in crate::db) fn read_u32(&mut self) -> Result<u32, ByteDecodeError> {
Ok(u32::from_be_bytes(self.read_array()?))
}
pub(in crate::db) fn read_u64(&mut self) -> Result<u64, ByteDecodeError> {
Ok(u64::from_be_bytes(self.read_array()?))
}
pub(in crate::db) fn read_i64(&mut self) -> Result<i64, ByteDecodeError> {
Ok(i64::from_be_bytes(self.read_array()?))
}
pub(in crate::db) fn read_i128(&mut self) -> Result<i128, ByteDecodeError> {
Ok(i128::from_be_bytes(self.read_array()?))
}
pub(in crate::db) fn read_u128(&mut self) -> Result<u128, ByteDecodeError> {
Ok(u128::from_be_bytes(self.read_array()?))
}
pub(in crate::db) fn read_u16_le(&mut self) -> Result<u16, ByteDecodeError> {
Ok(u16::from_le_bytes(self.read_array()?))
}
pub(in crate::db) fn read_u32_le(&mut self) -> Result<u32, ByteDecodeError> {
Ok(u32::from_le_bytes(self.read_array()?))
}
pub(in crate::db) fn read_u64_le(&mut self) -> Result<u64, ByteDecodeError> {
Ok(u64::from_le_bytes(self.read_array()?))
}
pub(in crate::db) fn read_u128_le(&mut self) -> Result<u128, ByteDecodeError> {
Ok(u128::from_le_bytes(self.read_array()?))
}
pub(in crate::db) fn read_len_prefixed_bytes_le(
&mut self,
) -> Result<&'a [u8], ByteDecodeError> {
let len = usize::try_from(self.read_u32_le()?).map_err(|_| ByteDecodeError)?;
self.read_exact(len)
}
pub(in crate::db) fn read_len_prefixed_bytes(&mut self) -> Result<&'a [u8], ByteDecodeError> {
self.read_bounded_len_prefixed_bytes(usize::MAX)
}
pub(in crate::db) fn read_bounded_len_prefixed_bytes(
&mut self,
max: usize,
) -> Result<&'a [u8], ByteDecodeError> {
let len = usize::try_from(self.read_u32()?).map_err(|_| ByteDecodeError)?;
if len > max {
return Err(ByteDecodeError);
}
self.read_exact(len)
}
pub(in crate::db) fn read_string(&mut self) -> Result<String, ByteDecodeError> {
self.read_bounded_string(usize::MAX)
}
pub(in crate::db) fn read_bounded_string(
&mut self,
max: usize,
) -> Result<String, ByteDecodeError> {
let bytes = self.read_bounded_len_prefixed_bytes(max)?;
std::str::from_utf8(bytes)
.map(str::to_string)
.map_err(|_| ByteDecodeError)
}
pub(in crate::db) const fn finish(self) -> Result<(), ByteDecodeError> {
if self.remaining() == 0 {
Ok(())
} else {
Err(ByteDecodeError)
}
}
}
#[cfg(test)]
mod tests {
use super::{ByteDecodeError, ByteReader};
#[test]
fn little_endian_reads_preserve_known_wire_bytes() {
assert_eq!(ByteReader::new(&[0x34, 0x12]).read_u16_le(), Ok(0x1234));
assert_eq!(
ByteReader::new(&[0x78, 0x56, 0x34, 0x12]).read_u32_le(),
Ok(0x1234_5678)
);
let bytes = [0xef, 0xcd, 0xab, 0x89, 0x67, 0x45, 0x23, 0x01];
assert_eq!(
ByteReader::new(&bytes).read_u64_le(),
Ok(0x0123_4567_89ab_cdef)
);
let bytes = [
0xef, 0xcd, 0xab, 0x89, 0x67, 0x45, 0x23, 0x01, 0, 0, 0, 0, 0, 0, 0, 0x80,
];
assert_eq!(
ByteReader::new(&bytes).read_u128_le(),
Ok(0x8000_0000_0000_0000_0123_4567_89ab_cdef)
);
let framed = [2, 0, 0, 0, b'a', b'b'];
let mut reader = ByteReader::new(&framed);
assert_eq!(reader.read_len_prefixed_bytes_le(), Ok(b"ab".as_slice()));
assert_eq!(reader.finish(), Ok(()));
for end in 0..framed.len() {
assert_eq!(
ByteReader::new(&framed[..end]).read_len_prefixed_bytes_le(),
Err(ByteDecodeError)
);
}
}
#[test]
fn failed_byte_reads_preserve_the_unread_payload() {
let mut reader = ByteReader::new(&[1, 2, 3]);
assert_eq!(reader.read_u8(), Ok(1));
assert_eq!(reader.read_exact(usize::MAX), Err(ByteDecodeError));
assert_eq!(reader.read_array::<3>(), Err(ByteDecodeError));
assert_eq!(reader.read_exact(2), Ok([2, 3].as_slice()));
assert_eq!(reader.finish(), Ok(()));
}
#[test]
fn length_and_text_boundaries_reject_malformed_payloads() {
let bytes = [0, 0, 0, 2, b'a', b'b'];
assert_eq!(
ByteReader::new(&bytes).read_bounded_string(2),
Ok("ab".into())
);
assert_eq!(
ByteReader::new(&bytes).read_bounded_string(1),
Err(ByteDecodeError)
);
for end in 0..bytes.len() {
assert_eq!(
ByteReader::new(&bytes[..end]).read_string(),
Err(ByteDecodeError)
);
}
assert_eq!(
ByteReader::new(&[0, 0, 0, 1, 0xff]).read_string(),
Err(ByteDecodeError)
);
assert_eq!(ByteReader::new(&bytes).finish(), Err(ByteDecodeError));
}
}