use crate::error::{ByteOffset, DecodeContext, DecodeError, DecodeErrorKind};
#[derive(Debug)]
pub struct Cursor<'a> {
data: &'a [u8],
pos: usize,
}
impl<'a> Cursor<'a> {
pub fn new(data: &'a [u8]) -> Self {
Self { data, pos: 0 }
}
pub fn position(&self) -> usize {
self.pos
}
pub fn remaining(&self) -> &'a [u8] {
&self.data[self.pos..]
}
pub fn capacity_hint(&self, count: u32) -> usize {
(count as usize).min(self.remaining().len())
}
pub fn original(&self) -> &'a [u8] {
self.data
}
pub fn is_empty(&self) -> bool {
self.pos >= self.data.len()
}
pub fn read_byte(&mut self) -> Result<u8, DecodeError> {
if self.pos < self.data.len() {
let b = self.data[self.pos];
self.pos += 1;
Ok(b)
} else {
Err(DecodeError {
offset: ByteOffset(self.pos),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::UnexpectedEof,
})
}
}
pub fn read_bytes(&mut self, n: usize) -> Result<&'a [u8], DecodeError> {
if self.pos + n <= self.data.len() {
let slice = &self.data[self.pos..self.pos + n];
self.pos += n;
Ok(slice)
} else {
Err(DecodeError {
offset: ByteOffset(self.pos),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::UnexpectedEof,
})
}
}
pub fn advance(&mut self, n: usize) -> Result<(), DecodeError> {
if self.pos + n <= self.data.len() {
self.pos += n;
Ok(())
} else {
Err(DecodeError {
offset: ByteOffset(self.pos),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::UnexpectedEof,
})
}
}
}
pub fn decode_u32(cursor: &mut Cursor<'_>) -> Result<u32, DecodeError> {
let start = cursor.position();
let mut result: u32 = 0;
let mut shift: u32 = 0;
for i in 0..5 {
let byte = cursor.read_byte().map_err(|mut e| {
e.context = DecodeContext::Leb128;
e.offset = ByteOffset(start);
e
})?;
let low_bits = u32::from(byte & 0x7F);
if i == 4 && (byte & 0xF0) != 0 {
return Err(DecodeError {
offset: ByteOffset(start),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::Leb128Overflow,
});
}
result |= low_bits << shift;
if byte & 0x80 == 0 {
return Ok(result);
}
shift += 7;
}
Err(DecodeError {
offset: ByteOffset(start),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::Leb128TooLong,
})
}
pub fn decode_u64(cursor: &mut Cursor<'_>) -> Result<u64, DecodeError> {
let start = cursor.position();
let mut result: u64 = 0;
let mut shift: u32 = 0;
for i in 0..10 {
let byte = cursor.read_byte().map_err(|mut e| {
e.context = DecodeContext::Leb128;
e.offset = ByteOffset(start);
e
})?;
let low_bits = u64::from(byte & 0x7F);
if i == 9 && (byte & 0xFE) != 0 {
return Err(DecodeError {
offset: ByteOffset(start),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::Leb128Overflow,
});
}
result |= low_bits << shift;
if byte & 0x80 == 0 {
return Ok(result);
}
shift += 7;
}
Err(DecodeError {
offset: ByteOffset(start),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::Leb128TooLong,
})
}
pub fn decode_i32(cursor: &mut Cursor<'_>) -> Result<i32, DecodeError> {
let start = cursor.position();
let mut result: i32 = 0;
let mut shift: u32 = 0;
for i in 0..5 {
let byte = cursor.read_byte().map_err(|mut e| {
e.context = DecodeContext::Leb128;
e.offset = ByteOffset(start);
e
})?;
let low_bits = i32::from(byte & 0x7F);
result |= low_bits << shift;
shift += 7;
if byte & 0x80 == 0 {
if i == 4 {
let sign = byte & 0x08;
let extension = byte & 0x70;
if (sign == 0 && extension != 0) || (sign != 0 && extension != 0x70) {
return Err(DecodeError {
offset: ByteOffset(start),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::Leb128Overflow,
});
}
} else if shift < 32 && (byte & 0x40) != 0 {
result |= !0 << shift;
}
return Ok(result);
}
}
Err(DecodeError {
offset: ByteOffset(start),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::Leb128TooLong,
})
}
pub fn decode_i64(cursor: &mut Cursor<'_>) -> Result<i64, DecodeError> {
let start = cursor.position();
let mut result: i64 = 0;
let mut shift: u32 = 0;
for i in 0..10 {
let byte = cursor.read_byte().map_err(|mut e| {
e.context = DecodeContext::Leb128;
e.offset = ByteOffset(start);
e
})?;
let low_bits = i64::from(byte & 0x7F);
result |= low_bits << shift;
shift += 7;
if byte & 0x80 == 0 {
if i == 9 {
let sign = byte & 0x01;
let extension = byte & 0x7E;
if (sign == 0 && extension != 0) || (sign != 0 && extension != 0x7E) {
return Err(DecodeError {
offset: ByteOffset(start),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::Leb128Overflow,
});
}
} else if shift < 64 && (byte & 0x40) != 0 {
result |= !0i64 << shift;
}
return Ok(result);
}
}
Err(DecodeError {
offset: ByteOffset(start),
context: DecodeContext::Leb128,
kind: DecodeErrorKind::Leb128TooLong,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn u32_zero() {
let mut c = Cursor::new(&[0x00]);
assert_eq!(decode_u32(&mut c).unwrap(), 0);
}
#[test]
fn u32_single_byte() {
let mut c = Cursor::new(&[0x08]);
assert_eq!(decode_u32(&mut c).unwrap(), 8);
}
#[test]
fn u32_max_single_byte() {
let mut c = Cursor::new(&[0x7F]);
assert_eq!(decode_u32(&mut c).unwrap(), 127);
}
#[test]
fn u32_two_bytes() {
let mut c = Cursor::new(&[0x80, 0x01]);
assert_eq!(decode_u32(&mut c).unwrap(), 128);
}
#[test]
fn u32_624485() {
let mut c = Cursor::new(&[0xE5, 0x8E, 0x26]);
assert_eq!(decode_u32(&mut c).unwrap(), 624485);
}
#[test]
fn u32_max_value() {
let mut c = Cursor::new(&[0xFF, 0xFF, 0xFF, 0xFF, 0x0F]);
assert_eq!(decode_u32(&mut c).unwrap(), u32::MAX);
}
#[test]
fn u32_overflow_fifth_byte() {
let mut c = Cursor::new(&[0xFF, 0xFF, 0xFF, 0xFF, 0x1F]);
let err = decode_u32(&mut c).unwrap_err();
assert_eq!(err.kind, DecodeErrorKind::Leb128Overflow);
}
#[test]
fn u32_too_long() {
let mut c = Cursor::new(&[0x80, 0x80, 0x80, 0x80, 0x80, 0x00]);
let err = decode_u32(&mut c).unwrap_err();
assert_eq!(err.kind, DecodeErrorKind::Leb128Overflow);
}
#[test]
fn u32_unexpected_eof() {
let mut c = Cursor::new(&[0x80]);
let err = decode_u32(&mut c).unwrap_err();
assert_eq!(err.kind, DecodeErrorKind::UnexpectedEof);
}
#[test]
fn i32_zero() {
let mut c = Cursor::new(&[0x00]);
assert_eq!(decode_i32(&mut c).unwrap(), 0);
}
#[test]
fn i32_positive() {
let mut c = Cursor::new(&[0x08]);
assert_eq!(decode_i32(&mut c).unwrap(), 8);
}
#[test]
fn i32_negative_one() {
let mut c = Cursor::new(&[0x7F]);
assert_eq!(decode_i32(&mut c).unwrap(), -1);
}
#[test]
fn i32_negative_two() {
let mut c = Cursor::new(&[0x7E]);
assert_eq!(decode_i32(&mut c).unwrap(), -2);
}
#[test]
fn i32_negative_128() {
let mut c = Cursor::new(&[0x80, 0x7F]);
assert_eq!(decode_i32(&mut c).unwrap(), -128);
}
#[test]
fn i32_min_value() {
let mut c = Cursor::new(&[0x80, 0x80, 0x80, 0x80, 0x78]);
assert_eq!(decode_i32(&mut c).unwrap(), i32::MIN);
}
#[test]
fn i32_max_value() {
let mut c = Cursor::new(&[0xFF, 0xFF, 0xFF, 0xFF, 0x07]);
assert_eq!(decode_i32(&mut c).unwrap(), i32::MAX);
}
#[test]
fn u64_zero() {
let mut c = Cursor::new(&[0x00]);
assert_eq!(decode_u64(&mut c).unwrap(), 0);
}
#[test]
fn u64_max_value() {
let mut c = Cursor::new(&[0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0x01]);
assert_eq!(decode_u64(&mut c).unwrap(), u64::MAX);
}
#[test]
fn u64_overflow() {
let mut c = Cursor::new(&[0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0x03]);
let err = decode_u64(&mut c).unwrap_err();
assert_eq!(err.kind, DecodeErrorKind::Leb128Overflow);
}
#[test]
fn i64_negative_one() {
let mut c = Cursor::new(&[0x7F]);
assert_eq!(decode_i64(&mut c).unwrap(), -1i64);
}
#[test]
fn i64_min_value() {
let mut c = Cursor::new(&[0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x80, 0x7F]);
assert_eq!(decode_i64(&mut c).unwrap(), i64::MIN);
}
#[test]
fn i64_max_value() {
let mut c = Cursor::new(&[0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0xFF, 0x00]);
assert_eq!(decode_i64(&mut c).unwrap(), i64::MAX);
}
#[test]
fn cursor_tracks_position() {
let mut c = Cursor::new(&[0x01, 0x02, 0x03]);
assert_eq!(c.position(), 0);
c.read_byte().unwrap();
assert_eq!(c.position(), 1);
c.read_bytes(2).unwrap();
assert_eq!(c.position(), 3);
assert!(c.is_empty());
}
}