use crate::{
Allocator, Chunk, CircularBuffer, GlobalAllocator, StaticVec, ValidationError, Vec, WasmStream,
};
const CONTINUATION_BIT: u8 = 0b10000000;
const INTEGER_BIT_FLAG: u8 = !CONTINUATION_BIT;
pub struct Reader<'wasm> {
stream: &'wasm mut dyn WasmStream,
chunk_used: usize,
next: Option<Chunk>,
buffer: CircularBuffer<u8, 64>,
full_offset: usize,
}
impl<'wasm> Reader<'wasm> {
pub fn new(stream: &'wasm mut dyn WasmStream) -> Self {
Self {
stream,
chunk_used: 0,
next: None,
buffer: CircularBuffer::new(),
full_offset: 0,
}
}
pub fn offset(&self) -> usize {
self.full_offset
}
fn fill_buffer(&mut self) -> Result<(), ValidationError> {
if !self.buffer.is_empty() {
return Ok(());
}
if let Some(ref chunk) = self.next {
let remaining = chunk.len() - self.chunk_used;
if remaining > 0 {
let to_copy = remaining.min(self.buffer.capacity());
for i in 0..to_copy {
self.buffer.push(chunk[self.chunk_used + i]);
}
self.chunk_used += to_copy;
return Ok(());
}
}
if let Some(mut chunk) = self.next.take() {
chunk.return_(self.stream);
}
self.next = self
.stream
.read()
.map_err(ValidationError::ReaderError)?
.map(|inner| inner.into());
self.chunk_used = 0;
if let Some(ref chunk) = self.next {
if chunk.is_empty() {
return Err(ValidationError::Eof);
}
let to_copy = chunk.len().min(self.buffer.capacity());
for i in 0..to_copy {
self.buffer.push(chunk[i]);
}
self.chunk_used = to_copy;
Ok(())
} else {
Err(ValidationError::Eof)
}
}
fn peek_u8(&mut self) -> Result<u8, ValidationError> {
if let Some(&byte) = self.buffer.front() {
return Ok(byte);
}
self.fill_buffer()?;
self.buffer.front().copied().ok_or(ValidationError::Eof)
}
pub fn read_u8(&mut self) -> Result<u8, ValidationError> {
let byte = self.peek_u8()?;
self.buffer.pop_front();
self.full_offset += 1;
Ok(byte)
}
pub fn expect_u8(&mut self, expected: u8) -> Result<(), ValidationError> {
let byte = self.peek_u8()?;
if byte == expected {
self.read_u8()?;
Ok(())
} else {
Err(ValidationError::ExpectedTerminal(expected))
}
}
pub fn strip_bytes<const N: usize>(&mut self) -> Result<[u8; N], ValidationError> {
let mut result = [0u8; N];
for item in result.iter_mut().take(N) {
*item = self.read_u8()?;
}
Ok(result)
}
pub fn read_u32(&mut self) -> Result<u32, ValidationError> {
const PADDING_IN_LAST_BYTE_BIT_MASK: u8 = 0b01110000;
let mut result: u32 = 0;
let byte = self.read_u8()?;
result |= u32::from(byte & INTEGER_BIT_FLAG);
if byte & CONTINUATION_BIT == 0 {
return Ok(result);
}
let byte = self.read_u8()?;
result |= u32::from(byte & INTEGER_BIT_FLAG) << 7;
if byte & CONTINUATION_BIT == 0 {
return Ok(result);
}
let byte = self.read_u8()?;
result |= u32::from(byte & INTEGER_BIT_FLAG) << 14;
if byte & CONTINUATION_BIT == 0 {
return Ok(result);
}
let byte = self.read_u8()?;
result |= u32::from(byte & INTEGER_BIT_FLAG) << 21;
if byte & CONTINUATION_BIT == 0 {
return Ok(result);
}
let byte = self.read_u8()?;
result |= u32::from(byte & INTEGER_BIT_FLAG) << 28;
let has_next_byte = byte & CONTINUATION_BIT > 0;
let padding_bits_are_not_zero = byte & PADDING_IN_LAST_BYTE_BIT_MASK > 0;
if has_next_byte || padding_bits_are_not_zero {
return Err(ValidationError::MalformedInteger);
}
Ok(result)
}
pub fn read_f64(&mut self) -> Result<u64, ValidationError> {
let bytes = self.strip_bytes::<8>()?;
Ok(u64::from_le_bytes(bytes))
}
pub fn read_i32(&mut self) -> Result<i32, ValidationError> {
const PADDING_IN_LAST_BYTE_BITMASK: u8 = 0b01110000;
const SIGN_IN_LAST_BYTE_BITFLAG: u8 = 0b00001000;
const NUM_BITS: u32 = 32;
let mut result: i32 = 0;
let byte = self.read_u8()?;
result |= i32::from(byte & INTEGER_BIT_FLAG);
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 7;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i32::from(byte & INTEGER_BIT_FLAG) << 7;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 14;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i32::from(byte & INTEGER_BIT_FLAG) << 14;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 21;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i32::from(byte & INTEGER_BIT_FLAG) << 21;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 28;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i32::from(byte & INTEGER_BIT_FLAG) << 28;
let has_next_byte = byte & CONTINUATION_BIT > 0;
if has_next_byte {
return Err(ValidationError::MalformedInteger);
}
const PADDING_AND_SIGN_BITMASK: u8 =
PADDING_IN_LAST_BYTE_BITMASK | SIGN_IN_LAST_BYTE_BITFLAG;
let number_of_ones_in_padding_and_sign_bits =
(byte & PADDING_AND_SIGN_BITMASK).count_ones();
let padding_bits_match_sign_bit = number_of_ones_in_padding_and_sign_bits
== PADDING_AND_SIGN_BITMASK.count_ones()
|| number_of_ones_in_padding_and_sign_bits == 0;
if !padding_bits_match_sign_bit {
return Err(ValidationError::MalformedInteger);
}
Ok(result)
}
pub fn read_f32(&mut self) -> Result<u32, ValidationError> {
let bytes = self.strip_bytes::<4>()?;
Ok(u32::from_le_bytes(bytes))
}
pub fn read_i64(&mut self) -> Result<i64, ValidationError> {
const PADDING_IN_LAST_BYTE_BITMASK: u8 = 0b01111110;
const SIGN_IN_LAST_BYTE_BITFLAG: u8 = 0b00000001;
const NUM_BITS: u32 = 64;
let mut result: i64 = 0;
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG);
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 7;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 7;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 14;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 14;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 21;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 21;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 28;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 28;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 35;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 35;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 42;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 42;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 49;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 49;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 56;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 56;
if byte & CONTINUATION_BIT == 0 {
const NUM_UNSPECIFIED_BITS: u32 = NUM_BITS - 63;
let sign_extended_result = (result << NUM_UNSPECIFIED_BITS) >> NUM_UNSPECIFIED_BITS;
return Ok(sign_extended_result);
}
let byte = self.read_u8()?;
result |= i64::from(byte & INTEGER_BIT_FLAG) << 63;
let has_next_byte = byte & CONTINUATION_BIT > 0;
if has_next_byte {
return Err(ValidationError::MalformedInteger);
}
const PADDING_AND_SIGN_BITMASK: u8 =
PADDING_IN_LAST_BYTE_BITMASK | SIGN_IN_LAST_BYTE_BITFLAG;
let number_of_ones_in_padding_and_sign_bits =
(byte & PADDING_AND_SIGN_BITMASK).count_ones();
let padding_bits_match_sign_bit = number_of_ones_in_padding_and_sign_bits
== PADDING_AND_SIGN_BITMASK.count_ones()
|| number_of_ones_in_padding_and_sign_bits == 0;
if !padding_bits_match_sign_bit {
return Err(ValidationError::MalformedInteger);
}
Ok(result)
}
pub fn skip(&mut self, len: usize) -> Result<(), ValidationError> {
for _ in 0..len {
self.read_u8()?;
}
Ok(())
}
pub fn read_vec<T, F>(&mut self, read_element: F) -> Result<Vec<T>, ValidationError>
where
T: 'wasm,
F: FnMut(&mut Self) -> Result<T, ValidationError>,
{
self.read_vec_in(GlobalAllocator, read_element)
}
pub fn read_vec_stack<const SIZE: usize, T>(
&mut self,
mut read_element: impl FnMut(&mut Self) -> Result<T, ValidationError>,
) -> Result<StaticVec<T, SIZE>, ValidationError>
where
T: 'wasm,
{
let len = self.read_u32()?;
if len as usize > SIZE {
return Err(ValidationError::VecTooLong);
}
let mut out = StaticVec::new();
for _ in 0..len {
out.push(read_element(self)?)?;
}
Ok(out)
}
pub fn read_vec_in<T, F, VA>(
&mut self,
alloc: VA,
mut read_element: F,
) -> Result<Vec<T, VA>, ValidationError>
where
T: 'wasm,
F: FnMut(&mut Self) -> Result<T, ValidationError>,
VA: Allocator,
{
let len = self.read_u32()?;
let mut out = Vec::new_in(alloc, len)?;
for _ in 0..len {
out.push(read_element(self)?);
}
Ok(out)
}
}
impl<'wasm> Drop for Reader<'wasm> {
fn drop(&mut self) {
if let Some(mut chunk) = self.next.take() {
chunk.return_(self.stream);
}
}
}
#[cfg(test)]
mod tests {
}