Skip to main content

burn_pack/
reader.rs

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
22/// Reader for loading burnpack containers.
23pub struct Reader {
24    metadata: Metadata,
25    source: Source,
26    /// Absolute byte offset where the (256-byte aligned) tensor data section starts.
27    data_offset: usize,
28}
29
30impl Reader {
31    /// Load a pack from an in-memory [`Bytes`] buffer.
32    ///
33    /// Loading is lazy: only the header and metadata are parsed here. The buffer is kept as-is
34    /// (no copy, no share) and is only turned into zero-copy [`Bytes::view`] windows when you
35    /// consume the reader with [`into_tensors`](Self::into_tensors). For large models, prefer
36    /// [`from_file`](Self::from_file), which keeps tensor data file-backed and lazy.
37    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    /// Load a pack from a file.
52    ///
53    /// Only the header and metadata are read up front; the whole file is wrapped in a single
54    /// lazy [`Bytes::from_file`] source, and each tensor's data is a [`Bytes::view`] window into
55    /// it, read from disk only when accessed. This integrates with the Burn ecosystem's file
56    /// allocation (pinned-memory staging) for fast file-to-GPU transfers.
57    ///
58    /// If `path` has no extension and does not exist as given, the canonical
59    /// [`crate::EXTENSION`] (`.bpk`) is appended.
60    #[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    /// Finish construction once the header, metadata, and data source are known.
91    ///
92    /// Centralizes the truncation check and the aligned data-section offset so both
93    /// [`from_bytes`](Self::from_bytes) and [`from_file`](Self::from_file) stay in sync.
94    /// `available` is the number of bytes the source can actually supply.
95    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    /// Consume the reader, returning all tensors in sorted (alphabetical) name order.
112    ///
113    /// Each tensor's bytes are a zero-copy [`Bytes::view`] window into the source — no per-tensor
114    /// copy. For an in-memory source the buffer is [shared](Bytes::shared) once here (a cheap
115    /// `Arc` move, never a data copy) to make those views available; for a file source the windows
116    /// are file-backed and read lazily on access. Consuming `self` lets us hand the source's
117    /// ownership to the views directly, so loading never reads or copies tensor data eagerly.
118    pub fn into_tensors(self) -> Result<Vec<Tensor>, Error> {
119        let Reader {
120            metadata,
121            source,
122            data_offset,
123        } = self;
124
125        // Make the source view-capable: a plain in-memory buffer has no zero-copy window until
126        // it's shared behind an `Arc`, whereas a file-backed source already windows lazily (and
127        // must NOT be shared, or every view would materialize the whole file).
128        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    /// The user-supplied key/value metadata stored alongside the tensors.
149    ///
150    /// For per-tensor info (dtype/shape/param id), use [`into_tensors`](Self::into_tensors) — for a
151    /// file-backed reader that does not read any tensor data until a tensor's bytes are accessed.
152    pub fn metadata(&self) -> &alloc::collections::BTreeMap<String, String> {
153        &self.metadata.metadata
154    }
155
156    /// The typed scalars stored alongside the tensors.
157    ///
158    /// Empty for files written before scalar support.
159    pub fn scalars(&self) -> &alloc::collections::BTreeMap<String, crate::Scalar> {
160        &self.metadata.scalars
161    }
162
163    /// The names of all tensors in the pack, in sorted (alphabetical) order.
164    pub fn tensor_names(&self) -> Vec<&str> {
165        self.metadata.tensors.keys().map(|n| n.as_str()).collect()
166    }
167
168    /// Read a single tensor's raw little-endian bytes by name (always copies).
169    ///
170    /// Returns [`Error::TensorNotFound`] if no tensor with that name exists.
171    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                // A file-backed view reads just this tensor's range from disk.
183                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
196/// Compute and validate the absolute `[start, end)` byte range of a tensor.
197fn 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
236/// Parse and validate a header from a buffer that starts with it.
237fn 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
254/// Deserialize the CBOR metadata and validate the tensor count.
255fn 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
268/// Ensure the available bytes can hold every tensor the metadata claims.
269fn 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
297/// Borrow a tensor's `[start, end)` range out of an in-memory buffer, bounds-checked.
298fn 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
308/// Build a [`Tensor`] entry from a descriptor + its data bytes.
309fn 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
331/// Where a [`Reader`] gets its tensor data from.
332///
333/// Both variants hold a single [`Bytes`] spanning the whole container, and tensors are carved
334/// out of it with zero-copy [`Bytes::view`] windows.
335enum Source {
336    /// The whole pack lives in memory.
337    Memory(Bytes),
338    /// The pack lives in a file; tensor data is read lazily via file-backed [`Bytes::view`].
339    #[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// Verifies the on-disk layout invariant (256-byte tensor alignment), which needs access to the
349// internal `TensorDescriptor` offsets. Public round-trip/error tests live in `tests/`.
350#[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        // Odd sizes force the writer to insert padding between tensors.
369        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}