1use super::base::{
2 Error, FORMAT_VERSION, HEADER_SIZE, Header, MAX_CBOR_RECURSION_DEPTH, MAX_METADATA_SIZE,
3 MAX_TENSOR_COUNT, MAX_TENSOR_SIZE, Metadata, TensorDescriptor, aligned_data_section_start,
4};
5use super::tensor::Tensor;
6use alloc::format;
7use alloc::string::{String, ToString};
8use alloc::vec::Vec;
9use burn_std::{Bytes, Shape};
10
11#[cfg(feature = "std")]
12use super::base::MAX_FILE_SIZE;
13#[cfg(feature = "std")]
14use alloc::vec;
15#[cfg(feature = "std")]
16use std::fs::File;
17#[cfg(feature = "std")]
18use std::io::Read;
19#[cfg(feature = "std")]
20use std::path::Path;
21
22pub struct Reader {
24 metadata: Metadata,
25 source: Source,
26 data_offset: usize,
28}
29
30impl Reader {
31 pub fn from_bytes(bytes: Bytes) -> Result<Self, Error> {
38 let header = read_header(&bytes)?;
39 let metadata_end = HEADER_SIZE
40 .checked_add(header.metadata_size as usize)
41 .ok_or(Error::InvalidHeader)?;
42 if bytes.len() < metadata_end {
43 return Err(Error::InvalidHeader);
44 }
45 let metadata = parse_metadata(&bytes[HEADER_SIZE..metadata_end])?;
46
47 let available = bytes.len();
48 Self::assemble(&header, metadata, Source::Memory(bytes), available)
49 }
50
51 #[cfg(feature = "std")]
61 pub fn from_file<P: AsRef<Path>>(path: P) -> Result<Self, Error> {
62 let path = path.as_ref();
63 let path = if path.extension().is_none() && !path.exists() {
64 path.with_extension(crate::EXTENSION)
65 } else {
66 path.to_path_buf()
67 };
68
69 let mut file = File::open(&path).map_err(io_err)?;
70
71 let file_size = file.metadata().map_err(io_err)?.len();
72 if file_size > MAX_FILE_SIZE {
73 return Err(Error::ValidationError(format!(
74 "File size {file_size} bytes exceeds maximum allowed size of {MAX_FILE_SIZE} bytes"
75 )));
76 }
77
78 let mut header_bytes = [0u8; HEADER_SIZE];
79 file.read_exact(&mut header_bytes).map_err(io_err)?;
80 let header = read_header(&header_bytes)?;
81
82 let mut metadata_bytes = vec![0u8; header.metadata_size as usize];
83 file.read_exact(&mut metadata_bytes).map_err(io_err)?;
84 let metadata = parse_metadata(&metadata_bytes)?;
85
86 let source = Source::File(Bytes::from_file(path.as_path(), file_size, 0));
87 Self::assemble(&header, metadata, source, file_size as usize)
88 }
89
90 fn assemble(
96 header: &Header,
97 metadata: Metadata,
98 source: Source,
99 available: usize,
100 ) -> Result<Self, Error> {
101 let metadata_end = HEADER_SIZE + header.metadata_size as usize;
102 validate_total_size(&metadata, metadata_end, available)?;
103
104 Ok(Self {
105 metadata,
106 source,
107 data_offset: aligned_data_section_start(header.metadata_size as usize),
108 })
109 }
110
111 pub fn into_tensors(self) -> Result<Vec<Tensor>, Error> {
119 let Reader {
120 metadata,
121 source,
122 data_offset,
123 } = self;
124
125 let source = match source {
129 Source::Memory(bytes) => bytes.shared(),
130 #[cfg(feature = "std")]
131 Source::File(bytes) => bytes,
132 };
133
134 let mut tensors = Vec::with_capacity(metadata.tensors.len());
135 for (name, descriptor) in &metadata.tensors {
136 let (start, end) = tensor_range(data_offset, name, descriptor)?;
137 let bytes = source.view(start, end).map_err(|_| {
138 Error::ValidationError(format!(
139 "Tensor '{name}' data range {start}..{end} could not be viewed (source is {} bytes)",
140 source.len()
141 ))
142 })?;
143 tensors.push(make_tensor(name, descriptor, bytes)?);
144 }
145 Ok(tensors)
146 }
147
148 pub fn metadata(&self) -> &alloc::collections::BTreeMap<String, String> {
153 &self.metadata.metadata
154 }
155
156 pub fn scalars(&self) -> &alloc::collections::BTreeMap<String, crate::Scalar> {
160 &self.metadata.scalars
161 }
162
163 pub fn tensor_names(&self) -> Vec<&str> {
165 self.metadata.tensors.keys().map(|n| n.as_str()).collect()
166 }
167
168 pub fn tensor_data(&self, name: &str) -> Result<Vec<u8>, Error> {
172 let descriptor = self
173 .metadata
174 .tensors
175 .get(name)
176 .ok_or_else(|| Error::TensorNotFound(name.to_string()))?;
177 let (start, end) = tensor_range(self.data_offset, name, descriptor)?;
178
179 match &self.source {
180 #[cfg(feature = "std")]
181 Source::File(bytes) => {
182 let view = bytes.view(start, end).map_err(|_| {
184 Error::ValidationError(format!(
185 "Tensor '{name}' data range {start}..{end} could not be viewed"
186 ))
187 })?;
188 let slice: &[u8] = &view;
189 Ok(slice.to_vec())
190 }
191 Source::Memory(bytes) => Ok(memory_chunk(bytes, start, end)?.to_vec()),
192 }
193 }
194}
195
196fn tensor_range(
198 data_offset: usize,
199 name: &str,
200 descriptor: &TensorDescriptor,
201) -> Result<(usize, usize), Error> {
202 let to_usize = |offset: u64| -> Result<usize, Error> {
203 offset.try_into().map_err(|_| {
204 Error::ValidationError(format!(
205 "Tensor '{name}' has corrupted offset data: offset {offset} exceeds platform maximum"
206 ))
207 })
208 };
209 let overflow = || {
210 Error::ValidationError(format!(
211 "Tensor '{name}' has corrupted offset data: overflow"
212 ))
213 };
214
215 let start = data_offset
216 .checked_add(to_usize(descriptor.data_offsets.0)?)
217 .ok_or_else(overflow)?;
218 let end = data_offset
219 .checked_add(to_usize(descriptor.data_offsets.1)?)
220 .ok_or_else(overflow)?;
221
222 if end < start {
223 return Err(Error::ValidationError(format!(
224 "Tensor '{name}' has corrupted offset data: end {end} < start {start}"
225 )));
226 }
227 if end - start > MAX_TENSOR_SIZE {
228 return Err(Error::ValidationError(format!(
229 "Tensor '{name}' size {} exceeds maximum allowed size of {MAX_TENSOR_SIZE} bytes (potential DoS attack)",
230 end - start
231 )));
232 }
233 Ok((start, end))
234}
235
236fn read_header(buf: &[u8]) -> Result<Header, Error> {
238 if buf.len() < HEADER_SIZE {
239 return Err(Error::InvalidHeader);
240 }
241 let header = Header::from_bytes(&buf[..HEADER_SIZE])?;
242 if header.version > FORMAT_VERSION {
243 return Err(Error::InvalidVersion);
244 }
245 if header.metadata_size > MAX_METADATA_SIZE {
246 return Err(Error::ValidationError(format!(
247 "Metadata size {} exceeds maximum allowed size of {MAX_METADATA_SIZE} bytes (potential DoS attack)",
248 header.metadata_size
249 )));
250 }
251 Ok(header)
252}
253
254fn parse_metadata(bytes: &[u8]) -> Result<Metadata, Error> {
256 let metadata: Metadata =
257 ciborium::de::from_reader_with_recursion_limit(bytes, MAX_CBOR_RECURSION_DEPTH)
258 .map_err(|e| Error::MetadataDeserializationError(e.to_string()))?;
259 if metadata.tensors.len() > MAX_TENSOR_COUNT {
260 return Err(Error::ValidationError(format!(
261 "File contains {} tensors, exceeding maximum of {MAX_TENSOR_COUNT} (potential DoS attack)",
262 metadata.tensors.len()
263 )));
264 }
265 Ok(metadata)
266}
267
268fn validate_total_size(
270 metadata: &Metadata,
271 metadata_end: usize,
272 available: usize,
273) -> Result<(), Error> {
274 if metadata.tensors.is_empty() {
275 return Ok(());
276 }
277 let max_offset = metadata
278 .tensors
279 .values()
280 .map(|t| t.data_offsets.1)
281 .max()
282 .unwrap_or(0);
283 let max_offset: usize = max_offset.try_into().map_err(|_| {
284 Error::ValidationError(format!("Data offset {max_offset} exceeds platform maximum"))
285 })?;
286 let min_size = metadata_end
287 .checked_add(max_offset)
288 .ok_or_else(|| Error::ValidationError("File size calculation overflow".into()))?;
289 if available < min_size {
290 return Err(Error::ValidationError(format!(
291 "File truncated: expected at least {min_size} bytes, got {available} bytes"
292 )));
293 }
294 Ok(())
295}
296
297fn memory_chunk(source: &Bytes, start: usize, end: usize) -> Result<&[u8], Error> {
299 let data: &[u8] = source;
300 data.get(start..end).ok_or_else(|| {
301 Error::ValidationError(format!(
302 "Tensor data range {start}..{end} is out of bounds (buffer is {} bytes)",
303 data.len()
304 ))
305 })
306}
307
308fn make_tensor(name: &str, descriptor: &TensorDescriptor, bytes: Bytes) -> Result<Tensor, Error> {
310 let shape = descriptor
311 .shape
312 .iter()
313 .map(|&s| {
314 s.try_into().map_err(|_| {
315 Error::ValidationError(format!(
316 "Tensor '{name}' has corrupted shape data: dimension {s} exceeds platform maximum"
317 ))
318 })
319 })
320 .collect::<Result<Vec<usize>, Error>>()?;
321
322 Ok(Tensor::new(
323 name.to_string(),
324 descriptor.dtype,
325 Shape::from(shape),
326 descriptor.param_id,
327 bytes,
328 ))
329}
330
331enum Source {
336 Memory(Bytes),
338 #[cfg(feature = "std")]
340 File(Bytes),
341}
342
343#[cfg(feature = "std")]
344fn io_err(e: std::io::Error) -> Error {
345 Error::IoError(e.to_string())
346}
347
348#[cfg(all(test, feature = "std"))]
351mod tests {
352 use super::*;
353 use crate::{TENSOR_ALIGNMENT, Tensor, Writer};
354 use burn_std::DType;
355
356 fn tensor(name: &str, elems: usize) -> Tensor {
357 Tensor::new(
358 name.to_string(),
359 DType::F32,
360 alloc::vec![elems],
361 None,
362 Bytes::from_bytes_vec(alloc::vec![0u8; elems * 4]),
363 )
364 }
365
366 #[test]
367 fn tensor_offsets_are_256_aligned() {
368 let packed = Writer::new(vec![tensor("a", 3), tensor("b", 1), tensor("c", 2)])
370 .into_bytes()
371 .unwrap();
372 let reader = Reader::from_bytes(packed).unwrap();
373
374 for (name, descriptor) in &reader.metadata.tensors {
375 assert_eq!(
376 descriptor.data_offsets.0 % TENSOR_ALIGNMENT,
377 0,
378 "tensor '{name}' start offset is not 256-aligned"
379 );
380 }
381 }
382}