use super::base::{
Error, FORMAT_VERSION, HEADER_SIZE, Header, MAX_CBOR_RECURSION_DEPTH, MAX_METADATA_SIZE,
MAX_TENSOR_COUNT, MAX_TENSOR_SIZE, Metadata, TensorDescriptor, aligned_data_section_start,
};
use super::tensor::Tensor;
use alloc::format;
use alloc::string::{String, ToString};
use alloc::vec::Vec;
use burn_std::{Bytes, Shape};
#[cfg(feature = "std")]
use super::base::MAX_FILE_SIZE;
#[cfg(feature = "std")]
use alloc::vec;
#[cfg(feature = "std")]
use std::fs::File;
#[cfg(feature = "std")]
use std::io::Read;
#[cfg(feature = "std")]
use std::path::Path;
pub struct Reader {
metadata: Metadata,
source: Source,
data_offset: usize,
}
impl Reader {
pub fn from_bytes(bytes: Bytes) -> Result<Self, Error> {
let header = read_header(&bytes)?;
let metadata_end = HEADER_SIZE
.checked_add(header.metadata_size as usize)
.ok_or(Error::InvalidHeader)?;
if bytes.len() < metadata_end {
return Err(Error::InvalidHeader);
}
let metadata = parse_metadata(&bytes[HEADER_SIZE..metadata_end])?;
let available = bytes.len();
Self::assemble(&header, metadata, Source::Memory(bytes), available)
}
#[cfg(feature = "std")]
pub fn from_file<P: AsRef<Path>>(path: P) -> Result<Self, Error> {
let path = path.as_ref();
let path = if path.extension().is_none() && !path.exists() {
path.with_extension(crate::EXTENSION)
} else {
path.to_path_buf()
};
let mut file = File::open(&path).map_err(io_err)?;
let file_size = file.metadata().map_err(io_err)?.len();
if file_size > MAX_FILE_SIZE {
return Err(Error::ValidationError(format!(
"File size {file_size} bytes exceeds maximum allowed size of {MAX_FILE_SIZE} bytes"
)));
}
let mut header_bytes = [0u8; HEADER_SIZE];
file.read_exact(&mut header_bytes).map_err(io_err)?;
let header = read_header(&header_bytes)?;
let mut metadata_bytes = vec![0u8; header.metadata_size as usize];
file.read_exact(&mut metadata_bytes).map_err(io_err)?;
let metadata = parse_metadata(&metadata_bytes)?;
let source = Source::File(Bytes::from_file(path.as_path(), file_size, 0));
Self::assemble(&header, metadata, source, file_size as usize)
}
fn assemble(
header: &Header,
metadata: Metadata,
source: Source,
available: usize,
) -> Result<Self, Error> {
let metadata_end = HEADER_SIZE + header.metadata_size as usize;
validate_total_size(&metadata, metadata_end, available)?;
Ok(Self {
metadata,
source,
data_offset: aligned_data_section_start(header.metadata_size as usize),
})
}
pub fn into_tensors(self) -> Result<Vec<Tensor>, Error> {
let Reader {
metadata,
source,
data_offset,
} = self;
let source = match source {
Source::Memory(bytes) => bytes.shared(),
#[cfg(feature = "std")]
Source::File(bytes) => bytes,
};
let mut tensors = Vec::with_capacity(metadata.tensors.len());
for (name, descriptor) in &metadata.tensors {
let (start, end) = tensor_range(data_offset, name, descriptor)?;
let bytes = source.view(start, end).map_err(|_| {
Error::ValidationError(format!(
"Tensor '{name}' data range {start}..{end} could not be viewed (source is {} bytes)",
source.len()
))
})?;
tensors.push(make_tensor(name, descriptor, bytes)?);
}
Ok(tensors)
}
pub fn metadata(&self) -> &alloc::collections::BTreeMap<String, String> {
&self.metadata.metadata
}
pub fn scalars(&self) -> &alloc::collections::BTreeMap<String, crate::Scalar> {
&self.metadata.scalars
}
pub fn tensor_names(&self) -> Vec<&str> {
self.metadata.tensors.keys().map(|n| n.as_str()).collect()
}
pub fn tensor_data(&self, name: &str) -> Result<Vec<u8>, Error> {
let descriptor = self
.metadata
.tensors
.get(name)
.ok_or_else(|| Error::TensorNotFound(name.to_string()))?;
let (start, end) = tensor_range(self.data_offset, name, descriptor)?;
match &self.source {
#[cfg(feature = "std")]
Source::File(bytes) => {
let view = bytes.view(start, end).map_err(|_| {
Error::ValidationError(format!(
"Tensor '{name}' data range {start}..{end} could not be viewed"
))
})?;
let slice: &[u8] = &view;
Ok(slice.to_vec())
}
Source::Memory(bytes) => Ok(memory_chunk(bytes, start, end)?.to_vec()),
}
}
}
fn tensor_range(
data_offset: usize,
name: &str,
descriptor: &TensorDescriptor,
) -> Result<(usize, usize), Error> {
let to_usize = |offset: u64| -> Result<usize, Error> {
offset.try_into().map_err(|_| {
Error::ValidationError(format!(
"Tensor '{name}' has corrupted offset data: offset {offset} exceeds platform maximum"
))
})
};
let overflow = || {
Error::ValidationError(format!(
"Tensor '{name}' has corrupted offset data: overflow"
))
};
let start = data_offset
.checked_add(to_usize(descriptor.data_offsets.0)?)
.ok_or_else(overflow)?;
let end = data_offset
.checked_add(to_usize(descriptor.data_offsets.1)?)
.ok_or_else(overflow)?;
if end < start {
return Err(Error::ValidationError(format!(
"Tensor '{name}' has corrupted offset data: end {end} < start {start}"
)));
}
if end - start > MAX_TENSOR_SIZE {
return Err(Error::ValidationError(format!(
"Tensor '{name}' size {} exceeds maximum allowed size of {MAX_TENSOR_SIZE} bytes (potential DoS attack)",
end - start
)));
}
Ok((start, end))
}
fn read_header(buf: &[u8]) -> Result<Header, Error> {
if buf.len() < HEADER_SIZE {
return Err(Error::InvalidHeader);
}
let header = Header::from_bytes(&buf[..HEADER_SIZE])?;
if header.version > FORMAT_VERSION {
return Err(Error::InvalidVersion);
}
if header.metadata_size > MAX_METADATA_SIZE {
return Err(Error::ValidationError(format!(
"Metadata size {} exceeds maximum allowed size of {MAX_METADATA_SIZE} bytes (potential DoS attack)",
header.metadata_size
)));
}
Ok(header)
}
fn parse_metadata(bytes: &[u8]) -> Result<Metadata, Error> {
let metadata: Metadata =
ciborium::de::from_reader_with_recursion_limit(bytes, MAX_CBOR_RECURSION_DEPTH)
.map_err(|e| Error::MetadataDeserializationError(e.to_string()))?;
if metadata.tensors.len() > MAX_TENSOR_COUNT {
return Err(Error::ValidationError(format!(
"File contains {} tensors, exceeding maximum of {MAX_TENSOR_COUNT} (potential DoS attack)",
metadata.tensors.len()
)));
}
Ok(metadata)
}
fn validate_total_size(
metadata: &Metadata,
metadata_end: usize,
available: usize,
) -> Result<(), Error> {
if metadata.tensors.is_empty() {
return Ok(());
}
let max_offset = metadata
.tensors
.values()
.map(|t| t.data_offsets.1)
.max()
.unwrap_or(0);
let max_offset: usize = max_offset.try_into().map_err(|_| {
Error::ValidationError(format!("Data offset {max_offset} exceeds platform maximum"))
})?;
let min_size = metadata_end
.checked_add(max_offset)
.ok_or_else(|| Error::ValidationError("File size calculation overflow".into()))?;
if available < min_size {
return Err(Error::ValidationError(format!(
"File truncated: expected at least {min_size} bytes, got {available} bytes"
)));
}
Ok(())
}
fn memory_chunk(source: &Bytes, start: usize, end: usize) -> Result<&[u8], Error> {
let data: &[u8] = source;
data.get(start..end).ok_or_else(|| {
Error::ValidationError(format!(
"Tensor data range {start}..{end} is out of bounds (buffer is {} bytes)",
data.len()
))
})
}
fn make_tensor(name: &str, descriptor: &TensorDescriptor, bytes: Bytes) -> Result<Tensor, Error> {
let shape = descriptor
.shape
.iter()
.map(|&s| {
s.try_into().map_err(|_| {
Error::ValidationError(format!(
"Tensor '{name}' has corrupted shape data: dimension {s} exceeds platform maximum"
))
})
})
.collect::<Result<Vec<usize>, Error>>()?;
Ok(Tensor::new(
name.to_string(),
descriptor.dtype,
Shape::from(shape),
descriptor.param_id,
bytes,
))
}
enum Source {
Memory(Bytes),
#[cfg(feature = "std")]
File(Bytes),
}
#[cfg(feature = "std")]
fn io_err(e: std::io::Error) -> Error {
Error::IoError(e.to_string())
}
#[cfg(all(test, feature = "std"))]
mod tests {
use super::*;
use crate::{TENSOR_ALIGNMENT, Tensor, Writer};
use burn_std::DType;
fn tensor(name: &str, elems: usize) -> Tensor {
Tensor::new(
name.to_string(),
DType::F32,
alloc::vec![elems],
None,
Bytes::from_bytes_vec(alloc::vec![0u8; elems * 4]),
)
}
#[test]
fn tensor_offsets_are_256_aligned() {
let packed = Writer::new(vec![tensor("a", 3), tensor("b", 1), tensor("c", 2)])
.into_bytes()
.unwrap();
let reader = Reader::from_bytes(packed).unwrap();
for (name, descriptor) in &reader.metadata.tensors {
assert_eq!(
descriptor.data_offsets.0 % TENSOR_ALIGNMENT,
0,
"tensor '{name}' start offset is not 256-aligned"
);
}
}
}