use crate::errors::AsMltError as _;
use crate::{Layer, MltError, MltResult, ParsedLayer};
const DEFAULT_MAX_BYTES: u32 = 20 * 1024 * 1024;
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct Decoder {
budget: MemBudget,
pub(crate) buffer_u32: Vec<u32>,
pub(crate) buffer_u64: Vec<u64>,
}
impl Decoder {
#[must_use]
pub fn with_max_size(max_bytes: u32) -> Self {
Self {
budget: MemBudget::with_max_size(max_bytes),
..Default::default()
}
}
pub fn decode_all<'a>(
&mut self,
layers: impl IntoIterator<Item = Layer<'a>>,
) -> MltResult<Vec<ParsedLayer<'a>>> {
layers
.into_iter()
.map(|l| l.decode_all(self))
.collect::<MltResult<_>>()
}
#[inline]
pub(crate) fn alloc<T>(&mut self, capacity: usize) -> MltResult<Vec<T>> {
let bytes = capacity.checked_mul(size_of::<T>()).or_overflow()?;
let bytes_u32 = u32::try_from(bytes).or_overflow()?;
self.budget.consume(bytes_u32)?;
Ok(Vec::with_capacity(capacity))
}
#[inline]
pub(crate) fn consume(&mut self, size: u32) -> MltResult<()> {
self.budget.consume(size)
}
#[inline]
pub(crate) fn consume_items<T>(&mut self, count: usize) -> MltResult<()> {
let bytes = count.checked_mul(size_of::<T>()).or_overflow()?;
self.budget.consume(u32::try_from(bytes).or_overflow()?)
}
#[inline]
pub(crate) fn adjust(&mut self, adjustment: u32) {
self.budget.adjust(adjustment);
}
#[inline]
pub(crate) fn adjust_alloc<T>(&mut self, buf: &[T], alloc_size: usize) -> MltResult<()> {
if buf.len() > alloc_size {
return Err(MltError::InvalidDecodingStreamSize(buf.len(), alloc_size));
}
let unused = (alloc_size - buf.len()) * size_of::<T>();
#[expect(
clippy::cast_possible_truncation,
reason = "unused <= alloc_size * size_of::<T>() which was verified to fit in u32 by alloc()"
)]
self.budget.adjust(unused as u32);
Ok(())
}
#[must_use]
pub fn consumed(&self) -> u32 {
self.budget.consumed()
}
pub fn reset_budget(&mut self) {
self.budget.reset();
}
}
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub struct Parser {
budget: MemBudget,
}
impl Parser {
#[must_use]
pub fn with_max_size(max_bytes: u32) -> Self {
Self {
budget: MemBudget::with_max_size(max_bytes),
}
}
pub fn parse_layers<'a>(&mut self, mut input: &'a [u8]) -> MltResult<Vec<Layer<'a>>> {
let mut result = Vec::new();
while !input.is_empty() {
let layer;
(input, layer) = Layer::from_bytes(input, self)?;
result.push(layer);
}
Ok(result)
}
#[inline]
pub(crate) fn reserve(&mut self, size: u32) -> MltResult<()> {
self.budget.consume(size)
}
#[must_use]
pub fn reserved(&self) -> u32 {
self.budget.consumed()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct MemBudget {
pub max_bytes: u32,
pub bytes_used: u32,
}
impl Default for MemBudget {
fn default() -> Self {
Self::with_max_size(DEFAULT_MAX_BYTES)
}
}
impl MemBudget {
#[must_use]
fn with_max_size(max_bytes: u32) -> Self {
Self {
max_bytes,
bytes_used: 0,
}
}
#[inline]
fn adjust(&mut self, adjustment: u32) {
self.bytes_used = self.bytes_used.checked_sub(adjustment).unwrap();
}
#[inline]
fn consume(&mut self, size: u32) -> MltResult<()> {
let accumulator = &mut self.bytes_used;
let max_bytes = self.max_bytes;
if let Some(new_value) = accumulator.checked_add(size).filter(|&v| v <= max_bytes) {
*accumulator = new_value;
Ok(())
} else {
Err(MltError::MemoryLimitExceeded {
limit: max_bytes,
used: *accumulator,
requested: size,
})
}
}
fn consumed(&self) -> u32 {
self.bytes_used
}
fn reset(&mut self) {
self.bytes_used = 0;
}
}