Skip to main content

hurray_io/file/
reader.rs

1use bytes::{Bytes, BytesMut};
2use hurray_core::{CompositeValidator, LayoutDescriptor, TensorDescriptor};
3use tokio::io::{AsyncRead, AsyncReadExt, AsyncSeek, AsyncSeekExt, SeekFrom};
4
5use crate::file::types::{IndexEntry, KvValue};
6use crate::file::{
7    file_flags, FILE_HEADER_SIZE, FILE_MAGIC, SUPPORTED_CONTAINER_VERSION_MAJOR, TRAILER_MAGIC,
8    TRAILER_SIZE,
9};
10use crate::{Error, Result};
11
12/// Default maximum composite nesting depth for [`FileReader::read_composite`].
13pub const DEFAULT_MAX_COMPOSITE_DEPTH: usize = 64;
14
15/// A tensor read from a Hurray file: descriptor plus zero-copy buffer views.
16#[derive(Debug)]
17pub struct FileTensor {
18    /// The tensor's name as recorded in the file index.
19    pub name: String,
20    /// The decoded tensor descriptor.
21    pub descriptor: TensorDescriptor,
22    /// Raw buffer bytes, one [`Bytes`] per buffer handle.
23    pub buffers: Vec<Bytes>,
24}
25
26/// One member of a composite read from a file: a plain tensor or a nested composite.
27#[derive(Debug)]
28pub enum FileItem {
29    /// A single (non-composite) tensor.
30    Tensor(FileTensor),
31    /// A nested composite.
32    Composite(FileComposite),
33}
34
35impl FileItem {
36    /// The item's governing descriptor: the tensor's descriptor, or the composite head.
37    pub fn descriptor(&self) -> &TensorDescriptor {
38        match self {
39            FileItem::Tensor(t) => &t.descriptor,
40            FileItem::Composite(c) => &c.head,
41        }
42    }
43}
44
45/// A composite tensor read from a file: its head plus its ordered members.
46///
47/// Membership was recovered from the head's `member_count` and file-offset adjacency, then
48/// validated with [`CompositeValidator`] (member count, partition coverage, overlay
49/// ordering). Members are ordered as written; each may itself be a composite (nesting).
50#[derive(Debug)]
51pub struct FileComposite {
52    /// The head's name, as recorded in the file index.
53    pub name: String,
54    /// The composite head descriptor (owns no data buffers).
55    pub head: TensorDescriptor,
56    /// The members, in write order. Each may itself be a composite.
57    pub members: Vec<FileItem>,
58}
59
60/// Reads tensors from a seekable Hurray file.
61///
62/// Open with [`FileReader::open`], then look up tensors by name with
63/// [`read_tensor`][FileReader::read_tensor] or inspect metadata with
64/// [`tensor_names`][FileReader::tensor_names] and [`kv`][FileReader::kv].
65///
66/// # Examples
67///
68/// ```no_run
69/// # #[tokio::main]
70/// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
71/// use hurray_io::file::FileReader;
72///
73/// let file = tokio::fs::File::open("model.hrry").await?;
74/// let mut reader = FileReader::open(file).await?;
75/// println!("tensors: {:?}", reader.tensor_names().collect::<Vec<_>>());
76/// let tensor = reader.read_tensor("embeddings").await?;
77/// println!("buffer[0]: {} bytes", tensor.buffers[0].len());
78/// # Ok(())
79/// # }
80/// ```
81#[derive(Debug)]
82pub struct FileReader<R> {
83    inner: R,
84    data_buffer_alignment: u64,
85    max_composite_depth: usize,
86    /// Footer index: one entry per tensor, in file-write order (or sorted if
87    /// `SORTED_INDEX` was set by the writer).
88    pub index: Vec<IndexEntry>,
89    /// File-level KV metadata (empty if the file has no KV section).
90    pub kv: Vec<(String, KvValue)>,
91}
92
93impl<R: AsyncRead + AsyncSeek + Unpin> FileReader<R> {
94    /// Opens a Hurray file: reads the header, trailer, index, and KV section.
95    ///
96    /// # Errors
97    ///
98    /// - [`Error::InvalidFileMagic`] — first 8 bytes are not `HRRYFILE`
99    /// - [`Error::UnsupportedContainerVersion`] — major version > 1
100    /// - [`Error::ReservedFileFlagBits`] — header `file_flags` has unknown bits set
101    /// - [`Error::InvalidTrailerMagic`] — trailer `trailer_magic` is not `HRRY`
102    /// - [`Error::IndexOverrunsTrailer`] — index section overlaps the trailer
103    /// - [`Error::IndexCrc32cMismatch`] — CRC-32C verification failed
104    /// - [`Error::DuplicateTensorName`] — index contains duplicate names
105    /// - [`Error::Io`] / [`Error::UnexpectedEof`] — I/O failures
106    pub async fn open(mut inner: R) -> Result<Self> {
107        // ── File header ───────────────────────────────────────────────────────
108        inner.seek(SeekFrom::Start(0)).await?;
109        let mut header = [0u8; 64];
110        inner
111            .read_exact(&mut header)
112            .await
113            .map_err(eof_to_unexpected)?;
114
115        if &header[0..8] != FILE_MAGIC {
116            return Err(Error::InvalidFileMagic);
117        }
118
119        let container_version_major = header[8];
120        if container_version_major > SUPPORTED_CONTAINER_VERSION_MAJOR {
121            return Err(Error::UnsupportedContainerVersion {
122                major: container_version_major,
123            });
124        }
125
126        let file_flags_val = u32::from_le_bytes(header[12..16].try_into().unwrap());
127        if file_flags_val & file_flags::RESERVED_MASK != 0 {
128            return Err(Error::ReservedFileFlagBits {
129                flags: file_flags_val,
130            });
131        }
132
133        let data_buffer_alignment = u32::from_le_bytes(header[16..20].try_into().unwrap()) as u64;
134
135        // ── Trailer ───────────────────────────────────────────────────────────
136        let file_size = inner.seek(SeekFrom::End(0)).await?;
137        if file_size < FILE_HEADER_SIZE + TRAILER_SIZE {
138            return Err(Error::InvalidHeader(
139                "file too small to contain a valid trailer".into(),
140            ));
141        }
142
143        inner
144            .seek(SeekFrom::Start(file_size - TRAILER_SIZE))
145            .await?;
146        let mut trailer = [0u8; 40];
147        inner
148            .read_exact(&mut trailer)
149            .await
150            .map_err(eof_to_unexpected)?;
151
152        if &trailer[36..40] != TRAILER_MAGIC {
153            return Err(Error::InvalidTrailerMagic);
154        }
155
156        let index_offset = u64::from_le_bytes(trailer[0..8].try_into().unwrap());
157        let index_length = u64::from_le_bytes(trailer[8..16].try_into().unwrap());
158        let kv_offset = u64::from_le_bytes(trailer[16..24].try_into().unwrap());
159        let kv_length = u32::from_le_bytes(trailer[24..28].try_into().unwrap());
160        let stored_crc = u32::from_le_bytes(trailer[28..32].try_into().unwrap());
161
162        if index_offset + index_length > file_size - TRAILER_SIZE {
163            return Err(Error::IndexOverrunsTrailer);
164        }
165
166        // ── Index section ────────────────────────────────────────────────────
167        inner.seek(SeekFrom::Start(index_offset)).await?;
168        let mut index_bytes = vec![0u8; index_length as usize];
169        inner
170            .read_exact(&mut index_bytes)
171            .await
172            .map_err(eof_to_unexpected)?;
173
174        if file_flags_val & file_flags::HAS_INDEX_CRC32C != 0 {
175            let computed = crc32c::crc32c(&index_bytes);
176            if computed != stored_crc {
177                return Err(Error::IndexCrc32cMismatch {
178                    stored: stored_crc,
179                    computed,
180                });
181            }
182        }
183
184        let index = parse_index(&index_bytes)?;
185
186        // ── KV section ───────────────────────────────────────────────────────
187        // Accept KV if either the flag or the trailer fields indicate its presence.
188        // A streaming writer may not set HAS_KV_METADATA in the header (no seek-back),
189        // so kv_offset != 0 is treated as authoritative.
190        let has_kv = (file_flags_val & file_flags::HAS_KV_METADATA != 0)
191            || (kv_offset != 0 && kv_length != 0);
192        let kv = if has_kv && kv_offset != 0 && kv_length != 0 {
193            inner.seek(SeekFrom::Start(kv_offset)).await?;
194            let mut kv_bytes = vec![0u8; kv_length as usize];
195            inner
196                .read_exact(&mut kv_bytes)
197                .await
198                .map_err(eof_to_unexpected)?;
199            parse_kv(&kv_bytes)?
200        } else {
201            Vec::new()
202        };
203
204        Ok(Self {
205            inner,
206            data_buffer_alignment,
207            max_composite_depth: DEFAULT_MAX_COMPOSITE_DEPTH,
208            index,
209            kv,
210        })
211    }
212
213    /// Sets the maximum composite nesting depth accepted by
214    /// [`read_composite`][FileReader::read_composite]. Default:
215    /// [`DEFAULT_MAX_COMPOSITE_DEPTH`].
216    pub fn with_max_composite_depth(mut self, max_composite_depth: usize) -> Self {
217        self.max_composite_depth = max_composite_depth;
218        self
219    }
220
221    /// Returns the names of all tensors in the file, in index order.
222    pub fn tensor_names(&self) -> impl Iterator<Item = &str> {
223        self.index.iter().map(|e| e.name.as_str())
224    }
225
226    /// Returns the file-level KV metadata.
227    pub fn kv(&self) -> &[(String, KvValue)] {
228        &self.kv
229    }
230
231    /// Decodes and returns the descriptor for `name` without reading buffer data.
232    ///
233    /// Useful for inspecting dtype, shape, or layout before deciding whether to
234    /// load the full tensor.
235    pub async fn read_descriptor(&mut self, name: &str) -> Result<TensorDescriptor> {
236        let entry = self.find_entry(name)?;
237        let offset = entry.descriptor_offset;
238        let len = entry.descriptor_length as usize;
239
240        self.inner.seek(SeekFrom::Start(offset)).await?;
241        let mut buf = vec![0u8; len];
242        self.inner
243            .read_exact(&mut buf)
244            .await
245            .map_err(eof_to_unexpected)?;
246        Ok(TensorDescriptor::decode(&buf)?)
247    }
248
249    /// Reads and returns a complete tensor (descriptor + all buffers).
250    pub async fn read_tensor(&mut self, name: &str) -> Result<FileTensor> {
251        let entry = self.find_entry(name)?;
252        let desc_offset = entry.descriptor_offset;
253        let desc_len = entry.descriptor_length as usize;
254        let data_offset = entry.data_offset;
255        let tensor_name = entry.name.clone();
256
257        // Decode descriptor
258        self.inner.seek(SeekFrom::Start(desc_offset)).await?;
259        let mut desc_buf = vec![0u8; desc_len];
260        self.inner
261            .read_exact(&mut desc_buf)
262            .await
263            .map_err(eof_to_unexpected)?;
264        let desc = TensorDescriptor::decode(&desc_buf)?;
265
266        // Read buffers from data region
267        self.inner.seek(SeekFrom::Start(data_offset)).await?;
268        let mut current = data_offset;
269        let n = desc.buffers.len();
270        let mut buffers = Vec::with_capacity(n);
271
272        for (i, handle) in desc.buffers.iter().enumerate() {
273            let byte_size = handle.byte_size() as usize;
274            let mut buf = BytesMut::with_capacity(byte_size);
275            buf.resize(byte_size, 0);
276            self.inner
277                .read_exact(&mut buf)
278                .await
279                .map_err(eof_to_unexpected)?;
280            current += byte_size as u64;
281            buffers.push(buf.freeze());
282
283            // Between consecutive buffers, seek past the alignment padding.
284            if i < n - 1 {
285                let aligned = align_up(current, self.data_buffer_alignment);
286                if aligned > current {
287                    self.inner.seek(SeekFrom::Start(aligned)).await?;
288                    current = aligned;
289                }
290            }
291        }
292
293        Ok(FileTensor {
294            name: tensor_name,
295            descriptor: desc,
296            buffers,
297        })
298    }
299
300    /// Reads a composite tensor by its head name, reassembling the head with its members.
301    ///
302    /// Members are recovered by **file-offset adjacency**: the head's declared
303    /// `member_count` tensors written immediately after it, recursively for nested
304    /// composites. Recovery keys on the descriptor offset, not index position, so it is
305    /// correct even when the file used the `SORTED_INDEX` option (which reorders the index
306    /// array but not the file layout). The reassembled group is validated with
307    /// [`CompositeValidator`].
308    ///
309    /// # Examples
310    ///
311    /// ```no_run
312    /// # #[tokio::main]
313    /// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
314    /// use hurray_io::file::FileReader;
315    ///
316    /// let file = tokio::fs::File::open("model.hrry").await?;
317    /// let mut reader = FileReader::open(file).await?;
318    /// let composite = reader.read_composite("weight").await?;
319    /// println!("{} member(s)", composite.members.len());
320    /// # Ok(())
321    /// # }
322    /// ```
323    ///
324    /// # Errors
325    ///
326    /// - [`Error::TensorNotFound`] — no tensor named `head_name`
327    /// - [`Error::NotAComposite`] — `head_name` exists but is not a composite head
328    /// - [`Error::TornComposite`] — fewer members follow the head than it declares
329    /// - [`Error::CompositeNestingTooDeep`] — nesting exceeded `max_composite_depth`
330    /// - [`Error::Core`] — composite validation failed
331    pub async fn read_composite(&mut self, head_name: &str) -> Result<FileComposite> {
332        // Offset-ordered snapshot of names: membership is recovered by write order (file
333        // offset), independent of the index array's order (which SORTED_INDEX permutes).
334        let mut ordered: Vec<(u64, String)> = self
335            .index
336            .iter()
337            .map(|e| (e.descriptor_offset, e.name.clone()))
338            .collect();
339        ordered.sort_unstable_by_key(|(offset, _)| *offset);
340
341        let start = ordered
342            .iter()
343            .position(|(_, name)| name == head_name)
344            .ok_or_else(|| Error::TensorNotFound(head_name.to_string()))?;
345
346        let names: Vec<String> = ordered.into_iter().map(|(_, name)| name).collect();
347        let mut cursor = start;
348        match self.consume_item(&names, &mut cursor, 0).await? {
349            FileItem::Composite(c) => Ok(c),
350            FileItem::Tensor(t) => Err(Error::NotAComposite(t.name)),
351        }
352    }
353
354    /// Consumes one item at `names[*cursor]` (advancing the cursor), recursing into members
355    /// when it is a composite head.
356    async fn consume_item(
357        &mut self,
358        names: &[String],
359        cursor: &mut usize,
360        depth: usize,
361    ) -> Result<FileItem> {
362        let name = names[*cursor].clone();
363        *cursor += 1;
364
365        let desc = self.read_descriptor(&name).await?;
366        let member_count = match &desc.layout {
367            LayoutDescriptor::Composite(c) => c.member_count,
368            _ => {
369                let tensor = self.read_tensor(&name).await?;
370                return Ok(FileItem::Tensor(tensor));
371            }
372        };
373
374        if depth >= self.max_composite_depth {
375            return Err(Error::CompositeNestingTooDeep {
376                limit: self.max_composite_depth,
377            });
378        }
379
380        let mut validator = CompositeValidator::new(&desc)?;
381        let mut members = Vec::with_capacity(member_count as usize);
382        for i in 0..member_count {
383            if *cursor >= names.len() {
384                return Err(Error::TornComposite {
385                    declared: member_count,
386                    actual: i,
387                });
388            }
389            // Box the recursive call: an async fn cannot name its own future.
390            let item = Box::pin(self.consume_item(names, cursor, depth + 1)).await?;
391            validator.push_member(item.descriptor())?;
392            members.push(item);
393        }
394        validator.finish()?;
395
396        Ok(FileItem::Composite(FileComposite {
397            name,
398            head: desc,
399            members,
400        }))
401    }
402
403    /// Returns the underlying reader.
404    pub fn into_inner(self) -> R {
405        self.inner
406    }
407
408    fn find_entry(&self, name: &str) -> Result<&IndexEntry> {
409        self.index
410            .iter()
411            .find(|e| e.name == name)
412            .ok_or_else(|| Error::TensorNotFound(name.to_string()))
413    }
414}
415
416// ── Parsing helpers ───────────────────────────────────────────────────────────
417
418fn align_up(offset: u64, alignment: u64) -> u64 {
419    let rem = offset % alignment;
420    if rem == 0 {
421        offset
422    } else {
423        offset + (alignment - rem)
424    }
425}
426
427fn eof_to_unexpected(e: std::io::Error) -> Error {
428    if e.kind() == std::io::ErrorKind::UnexpectedEof {
429        Error::UnexpectedEof
430    } else {
431        Error::Io(e)
432    }
433}
434
435fn read_u16(bytes: &[u8], pos: &mut usize) -> Result<u16> {
436    if *pos + 2 > bytes.len() {
437        return Err(Error::InvalidHeader("truncated field (u16)".into()));
438    }
439    let v = u16::from_le_bytes(bytes[*pos..*pos + 2].try_into().unwrap());
440    *pos += 2;
441    Ok(v)
442}
443
444fn read_u32(bytes: &[u8], pos: &mut usize) -> Result<u32> {
445    if *pos + 4 > bytes.len() {
446        return Err(Error::InvalidHeader("truncated field (u32)".into()));
447    }
448    let v = u32::from_le_bytes(bytes[*pos..*pos + 4].try_into().unwrap());
449    *pos += 4;
450    Ok(v)
451}
452
453fn read_u64(bytes: &[u8], pos: &mut usize) -> Result<u64> {
454    if *pos + 8 > bytes.len() {
455        return Err(Error::InvalidHeader("truncated field (u64)".into()));
456    }
457    let v = u64::from_le_bytes(bytes[*pos..*pos + 8].try_into().unwrap());
458    *pos += 8;
459    Ok(v)
460}
461
462fn read_bytes_slice<'a>(bytes: &'a [u8], pos: &mut usize, len: usize) -> Result<&'a [u8]> {
463    if *pos + len > bytes.len() {
464        return Err(Error::InvalidHeader("truncated byte slice".into()));
465    }
466    let s = &bytes[*pos..*pos + len];
467    *pos += len;
468    Ok(s)
469}
470
471fn parse_index(bytes: &[u8]) -> Result<Vec<IndexEntry>> {
472    let mut pos = 0usize;
473    let count = read_u64(bytes, &mut pos)? as usize;
474    let mut entries = Vec::with_capacity(count);
475    let mut seen = std::collections::HashSet::new();
476
477    for _ in 0..count {
478        let name_len = read_u16(bytes, &mut pos)? as usize;
479        if name_len == 0 {
480            return Err(Error::TensorNameEmpty);
481        }
482        let name_bytes = read_bytes_slice(bytes, &mut pos, name_len)?;
483        let name = std::str::from_utf8(name_bytes)
484            .map_err(|_| Error::InvalidHeader("tensor name is not valid UTF-8".into()))?
485            .to_string();
486        if !seen.insert(name.clone()) {
487            return Err(Error::DuplicateTensorName(name));
488        }
489
490        let descriptor_offset = read_u64(bytes, &mut pos)?;
491        let descriptor_length = read_u32(bytes, &mut pos)?;
492        let data_offset = read_u64(bytes, &mut pos)?;
493        let data_length = read_u64(bytes, &mut pos)?;
494        let _flags = read_u32(bytes, &mut pos)?; // reserved
495
496        entries.push(IndexEntry {
497            name,
498            descriptor_offset,
499            descriptor_length,
500            data_offset,
501            data_length,
502        });
503    }
504
505    Ok(entries)
506}
507
508fn parse_kv(bytes: &[u8]) -> Result<Vec<(String, KvValue)>> {
509    let mut pos = 0usize;
510    let count = read_u32(bytes, &mut pos)? as usize;
511    let mut entries = Vec::with_capacity(count);
512    let mut seen = std::collections::HashSet::new();
513
514    for _ in 0..count {
515        let key_len = read_u16(bytes, &mut pos)? as usize;
516        if key_len == 0 {
517            return Err(Error::KvKeyEmpty);
518        }
519        let key_bytes = read_bytes_slice(bytes, &mut pos, key_len)?;
520        let key = std::str::from_utf8(key_bytes)
521            .map_err(|_| Error::InvalidHeader("KV key is not valid UTF-8".into()))?
522            .to_string();
523        if !seen.insert(key.clone()) {
524            return Err(Error::DuplicateKvKey(key));
525        }
526
527        let value_tag = if pos < bytes.len() {
528            let t = bytes[pos];
529            pos += 1;
530            t
531        } else {
532            return Err(Error::InvalidHeader("KV value tag missing".into()));
533        };
534
535        let value = parse_kv_value(value_tag, bytes, &mut pos)?;
536        entries.push((key, value));
537    }
538
539    Ok(entries)
540}
541
542fn parse_kv_value(tag: u8, bytes: &[u8], pos: &mut usize) -> Result<KvValue> {
543    match tag {
544        0x01 => {
545            let len = read_u32(bytes, pos)? as usize;
546            let s = read_bytes_slice(bytes, pos, len)?;
547            let s = std::str::from_utf8(s)
548                .map_err(|_| Error::InvalidHeader("KV string is not valid UTF-8".into()))?
549                .to_string();
550            Ok(KvValue::String(s))
551        }
552        0x02 => Ok(KvValue::Int64(read_u64(bytes, pos)? as i64)),
553        0x03 => Ok(KvValue::Uint64(read_u64(bytes, pos)?)),
554        0x04 => {
555            let raw = read_u64(bytes, pos)?;
556            Ok(KvValue::Float64(f64::from_le_bytes(raw.to_le_bytes())))
557        }
558        0x05 => {
559            let b = read_bytes_slice(bytes, pos, 1)?[0];
560            match b {
561                0x00 => Ok(KvValue::Bool(false)),
562                0x01 => Ok(KvValue::Bool(true)),
563                _ => Err(Error::InvalidHeader(format!(
564                    "invalid KV bool byte 0x{b:02X}"
565                ))),
566            }
567        }
568        0x06 => {
569            let len = read_u32(bytes, pos)? as usize;
570            let v = read_bytes_slice(bytes, pos, len)?.to_vec();
571            Ok(KvValue::Bytes(v))
572        }
573        0x07 => {
574            let elem_tag = read_bytes_slice(bytes, pos, 1)?[0];
575            if elem_tag == 0x07 || elem_tag > 0x06 {
576                return Err(Error::InvalidHeader(format!(
577                    "invalid KV array element tag 0x{elem_tag:02X}"
578                )));
579            }
580            let count = read_u32(bytes, pos)? as usize;
581            let mut elements = Vec::with_capacity(count);
582            for _ in 0..count {
583                elements.push(parse_kv_value(elem_tag, bytes, pos)?);
584            }
585            Ok(KvValue::Array(elements))
586        }
587        _ => Err(Error::InvalidHeader(format!(
588            "unknown KV value tag 0x{tag:02X}"
589        ))),
590    }
591}