Skip to main content

hurray_io/file/
writer.rs

1use std::collections::HashSet;
2
3use hurray_core::{CompositeValidator, TensorDescriptor};
4use tokio::io::{AsyncWrite, AsyncWriteExt};
5
6use crate::file::types::{FileWriterOptions, KvValue};
7use crate::file::{file_flags, FILE_HEADER_SIZE, FILE_MAGIC, TRAILER_MAGIC};
8use crate::{Error, Result};
9
10/// A named node to write as part of a composite via [`FileWriter::write_composite`]:
11/// either a plain tensor or a nested composite. Every node carries a name because every
12/// tensor — head and members alike — gets its own footer-index entry (ADR-027 § Binding).
13pub enum FileCompositeNode<'a> {
14    /// A single (non-composite) tensor member.
15    Tensor {
16        /// Unique tensor name for this member's index entry.
17        name: &'a str,
18        /// The member's tensor descriptor.
19        descriptor: &'a TensorDescriptor,
20        /// One byte-slice per buffer handle in `descriptor.buffers`.
21        buffers: &'a [&'a [u8]],
22    },
23    /// A nested composite member: its head plus its own members.
24    Composite {
25        /// Unique tensor name for the nested head's index entry.
26        name: &'a str,
27        /// The nested composite's head descriptor.
28        head: &'a TensorDescriptor,
29        /// The nested composite's members, in order.
30        members: &'a [FileCompositeNode<'a>],
31    },
32}
33
34impl FileCompositeNode<'_> {
35    /// The node's governing descriptor: the tensor's descriptor, or the nested head.
36    fn descriptor(&self) -> &TensorDescriptor {
37        match self {
38            FileCompositeNode::Tensor { descriptor, .. } => descriptor,
39            FileCompositeNode::Composite { head, .. } => head,
40        }
41    }
42}
43
44/// Recursively validates a composite node tree, writing nothing. Reuses
45/// [`CompositeValidator`] at every level so [`FileWriter::write_composite`] can reject a
46/// torn/invalid composite before emitting any tensor.
47fn validate_composite_node(
48    head: &TensorDescriptor,
49    members: &[FileCompositeNode<'_>],
50) -> Result<()> {
51    let mut validator = CompositeValidator::new(head)?;
52    for member in members {
53        validator.push_member(member.descriptor())?;
54    }
55    validator.finish()?;
56
57    for member in members {
58        if let FileCompositeNode::Composite {
59            head: nested_head,
60            members: nested_members,
61            ..
62        } = member
63        {
64            validate_composite_node(nested_head, nested_members)?;
65        }
66    }
67    Ok(())
68}
69
70struct InternalEntry {
71    name: String,
72    descriptor_offset: u64,
73    descriptor_length: u32,
74    data_offset: u64,
75    data_length: u64,
76}
77
78/// Writes a Hurray file in a single forward pass without seeks.
79///
80/// The file format is:
81/// ```text
82/// [ File header      ]  64 bytes
83/// [ Tensor region    ]  descriptor → pad → buffers → pad  (repeated)
84/// [ KV section       ]  optional, written by finish()
85/// [ Index section    ]  written by finish()
86/// [ Trailer          ]  40 bytes
87/// ```
88///
89/// # Note on `HAS_KV_METADATA`
90///
91/// This writer sets `HAS_KV_METADATA = 0` in the file header because KV
92/// content is not known until [`finish`][FileWriter::finish] is called and a
93/// streaming writer cannot seek back to patch the header. The trailer's
94/// `kv_offset` field is the canonical indicator of KV presence; [`FileReader`]
95/// uses that field rather than the header flag.
96///
97/// [`FileReader`]: crate::file::FileReader
98///
99/// # Examples
100///
101/// ```no_run
102/// # #[tokio::main]
103/// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
104/// use hurray_core::{
105///     BufferHandle, DeviceTag, ElementType, LayoutDescriptor, Shape,
106///     SyncMode, TensorDescriptor, MIN_BUFFER_ALIGNMENT,
107/// };
108/// use hurray_io::file::{FileWriter, KvValue};
109///
110/// let handle = BufferHandle::new(64, MIN_BUFFER_ALIGNMENT, DeviceTag::Cpu, SyncMode::ProducerSynced)?;
111/// let shape = Shape::new(vec![4u64, 4]).unwrap();
112/// let desc = TensorDescriptor::new(
113///     1, 0, ElementType::Float32, shape, 0,
114///     LayoutDescriptor::RowMajor, vec![handle], None, None, None, None,
115/// )?;
116/// let data = vec![0u8; 64];
117///
118/// let file = tokio::fs::File::create("model.hrry").await?;
119/// let mut writer = FileWriter::new(file).await?;
120/// writer.write_tensor("embeddings", &desc, &[&data]).await?;
121/// writer.finish(vec![
122///     ("model".to_string(), KvValue::String("llama-3".to_string())),
123/// ]).await?;
124/// # Ok(())
125/// # }
126/// ```
127pub struct FileWriter<W> {
128    inner: W,
129    alignment: u64,
130    sorted_index: bool,
131    current_offset: u64,
132    entries: Vec<InternalEntry>,
133    seen_names: HashSet<String>,
134}
135
136static ZEROS: [u8; 4096] = [0u8; 4096];
137
138async fn write_zeroes<W: AsyncWrite + Unpin>(w: &mut W, mut count: u64) -> Result<()> {
139    while count > 0 {
140        let n = count.min(ZEROS.len() as u64) as usize;
141        w.write_all(&ZEROS[..n]).await?;
142        count -= n as u64;
143    }
144    Ok(())
145}
146
147fn pad_to(offset: u64, alignment: u64) -> u64 {
148    let rem = offset % alignment;
149    if rem == 0 {
150        0
151    } else {
152        alignment - rem
153    }
154}
155
156impl<W: AsyncWrite + Unpin> FileWriter<W> {
157    /// Creates a writer with default options (4096-byte buffer alignment, unsorted index).
158    pub async fn new(inner: W) -> Result<Self> {
159        Self::with_options(inner, FileWriterOptions::default()).await
160    }
161
162    /// Creates a writer with custom options.
163    pub async fn with_options(mut inner: W, options: FileWriterOptions) -> Result<Self> {
164        options.validate().map_err(Error::InvalidHeader)?;
165
166        let align = options.data_buffer_alignment as u64;
167        let sorted_index = options.sorted_index;
168
169        // File flags known at write time. HAS_KV_METADATA is not set here
170        // because KV content is supplied to finish(); see type-level note.
171        let mut flags: u32 = file_flags::HAS_INDEX_CRC32C;
172        if sorted_index {
173            flags |= file_flags::SORTED_INDEX;
174        }
175
176        // 64-byte file header
177        let mut header = [0u8; 64];
178        header[0..8].copy_from_slice(FILE_MAGIC);
179        header[8] = 1; // container_version_major
180        header[9] = 0; // container_version_minor
181                       // [10..12] _reserved = 0
182        header[12..16].copy_from_slice(&flags.to_le_bytes());
183        header[16..20].copy_from_slice(&(align as u32).to_le_bytes());
184        // first_descriptor_offset = 64 (immediately after the header)
185        header[20..28].copy_from_slice(&FILE_HEADER_SIZE.to_le_bytes());
186        // tensor_count_hint = UINT64_MAX (unknown at write time)
187        header[28..36].copy_from_slice(&u64::MAX.to_le_bytes());
188        // [36..64] _reserved_header = 0
189
190        inner.write_all(&header).await?;
191
192        Ok(Self {
193            inner,
194            alignment: align,
195            sorted_index,
196            current_offset: FILE_HEADER_SIZE,
197            entries: Vec::new(),
198            seen_names: HashSet::new(),
199        })
200    }
201
202    /// Encodes and writes one tensor.
203    ///
204    /// # Errors
205    ///
206    /// - [`Error::TensorNameEmpty`] / [`Error::TensorNameTooLong`] — invalid name
207    /// - [`Error::DuplicateTensorName`] — name already written to this file
208    /// - [`Error::MultiBufferLengthMismatch`] — wrong buffer count
209    /// - [`Error::BufferSizeMismatch`] — buffer length ≠ handle `byte_size`
210    /// - [`Error::Core`] — descriptor encoding failed
211    /// - [`Error::Io`] — underlying write error
212    pub async fn write_tensor(
213        &mut self,
214        name: &str,
215        desc: &TensorDescriptor,
216        buffers: &[&[u8]],
217    ) -> Result<()> {
218        if name.is_empty() {
219            return Err(Error::TensorNameEmpty);
220        }
221        if name.len() > 65535 {
222            return Err(Error::TensorNameTooLong { len: name.len() });
223        }
224        if !self.seen_names.insert(name.to_string()) {
225            return Err(Error::DuplicateTensorName(name.to_string()));
226        }
227
228        if buffers.len() != desc.buffers.len() {
229            return Err(Error::MultiBufferLengthMismatch {
230                declared: desc.buffers.len(),
231                actual: buffers.len(),
232            });
233        }
234        for (i, (handle, buf)) in desc.buffers.iter().zip(buffers).enumerate() {
235            let declared = handle.byte_size();
236            let actual = buf.len() as u64;
237            if declared != actual {
238                return Err(Error::BufferSizeMismatch {
239                    index: i,
240                    declared,
241                    actual,
242                });
243            }
244        }
245
246        let descriptor_offset = self.current_offset;
247        let encoded = desc.encode()?;
248        self.inner.write_all(&encoded).await?;
249        self.current_offset += encoded.len() as u64;
250        let descriptor_length = encoded.len() as u32;
251
252        // Pad descriptor end → first data-buffer alignment boundary
253        let pad = pad_to(self.current_offset, self.alignment);
254        write_zeroes(&mut self.inner, pad).await?;
255        self.current_offset += pad;
256
257        let data_offset = self.current_offset;
258        let n = buffers.len();
259
260        for (i, buf) in buffers.iter().enumerate() {
261            self.inner.write_all(buf).await?;
262            self.current_offset += buf.len() as u64;
263
264            // Pad between consecutive buffers (not after the last one)
265            if i < n - 1 {
266                let pad = pad_to(self.current_offset, self.alignment);
267                write_zeroes(&mut self.inner, pad).await?;
268                self.current_offset += pad;
269            }
270        }
271
272        // data_length covers bytes from data_offset through the last byte of the
273        // last buffer, excluding any trailing alignment padding for the next descriptor.
274        let data_length = self.current_offset - data_offset;
275
276        // Pad to 8-byte boundary so the next descriptor starts correctly.
277        let pad = pad_to(self.current_offset, 8);
278        write_zeroes(&mut self.inner, pad).await?;
279        self.current_offset += pad;
280
281        self.entries.push(InternalEntry {
282            name: name.to_string(),
283            descriptor_offset,
284            descriptor_length,
285            data_offset,
286            data_length,
287        });
288
289        Ok(())
290    }
291
292    /// Writes a composite tensor: its head, then every member's descriptor and data,
293    /// contiguously and in order (ADR-027 § Binding).
294    ///
295    /// Every tensor — the head and each member — gets its own footer-index entry, so all
296    /// are individually addressable by name via [`read_tensor`][crate::file::FileReader::read_tensor].
297    /// Membership is recoverable by [`read_composite`][crate::file::FileReader::read_composite]
298    /// from the head's `member_count` plus file-offset adjacency (the members are the tensors
299    /// written immediately after the head). Nested composites are written recursively.
300    ///
301    /// The whole group is validated up front — reusing [`CompositeValidator`] for member
302    /// count and per-rule constraints (partition exact-cover, overlay ordering) — before any
303    /// tensor is written.
304    ///
305    /// # Errors
306    ///
307    /// - [`Error::Core`] — the head is not a valid composite head, or validation failed
308    /// - the name/buffer errors of [`write_tensor`][FileWriter::write_tensor] for the head
309    ///   or any member
310    /// - [`Error::Io`] — underlying write error
311    ///
312    /// # Examples
313    ///
314    /// ```no_run
315    /// # #[tokio::main]
316    /// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
317    /// use hurray_core::{
318    ///     layout::{CompositeLayout, CompositionRule, LayoutDescriptor},
319    ///     ElementType, Shape, TensorDescriptor,
320    /// };
321    /// use hurray_io::file::{FileCompositeNode, FileWriter};
322    ///
323    /// let head = TensorDescriptor::new(
324    ///     1, 0, ElementType::Float32, Shape::new(vec![8u64, 8]).unwrap(), 0,
325    ///     LayoutDescriptor::Composite(CompositeLayout::new(CompositionRule::Partition, 2).unwrap()),
326    ///     vec![], None, None, None, None,
327    /// )?;
328    /// # let members: Vec<FileCompositeNode> = vec![];
329    /// let file = tokio::fs::File::create("model.hrry").await?;
330    /// let mut writer = FileWriter::new(file).await?;
331    /// writer.write_composite("weight", &head, &members).await?;
332    /// writer.finish(vec![]).await?;
333    /// # Ok(())
334    /// # }
335    /// ```
336    pub async fn write_composite(
337        &mut self,
338        head_name: &str,
339        head: &TensorDescriptor,
340        members: &[FileCompositeNode<'_>],
341    ) -> Result<()> {
342        // Validate the whole tree before writing any tensor.
343        validate_composite_node(head, members)?;
344
345        // Head first (no data buffers), then each member in order.
346        self.write_tensor(head_name, head, &[]).await?;
347        self.write_composite_members(members).await
348    }
349
350    /// Writes each member (recursively for nested composites) after the head.
351    async fn write_composite_members(&mut self, members: &[FileCompositeNode<'_>]) -> Result<()> {
352        for member in members {
353            match member {
354                FileCompositeNode::Tensor {
355                    name,
356                    descriptor,
357                    buffers,
358                } => {
359                    self.write_tensor(name, descriptor, buffers).await?;
360                }
361                FileCompositeNode::Composite {
362                    name,
363                    head,
364                    members,
365                } => {
366                    self.write_tensor(name, head, &[]).await?;
367                    // Box the recursive call: an async fn cannot name its own future.
368                    Box::pin(self.write_composite_members(members)).await?;
369                }
370            }
371        }
372        Ok(())
373    }
374
375    /// Writes the KV section, footer index, and trailer, then flushes.
376    ///
377    /// `kv` is a list of `(key, value)` pairs. Keys must be non-empty, at most
378    /// 65 535 bytes, and unique within the list.
379    pub async fn finish(mut self, kv: Vec<(String, KvValue)>) -> Result<W> {
380        validate_kv_keys(&kv)?;
381
382        // The tensor region is already 8-aligned after the last write_tensor().
383        // No additional padding needed before the KV section.
384
385        let (kv_offset, kv_length) = if kv.is_empty() {
386            (0u64, 0u32)
387        } else {
388            let kv_bytes = encode_kv_section(&kv)?;
389            let kv_offset = self.current_offset;
390            self.inner.write_all(&kv_bytes).await?;
391            self.current_offset += kv_bytes.len() as u64;
392
393            // Pad to 8-byte boundary for the index section
394            let pad = pad_to(self.current_offset, 8);
395            write_zeroes(&mut self.inner, pad).await?;
396            self.current_offset += pad;
397
398            (kv_offset, kv_bytes.len() as u32)
399        };
400
401        if self.sorted_index {
402            self.entries
403                .sort_unstable_by(|a, b| a.name.as_bytes().cmp(b.name.as_bytes()));
404        }
405
406        let index_bytes = encode_index_section(&self.entries);
407        let index_offset = self.current_offset;
408        let index_length = index_bytes.len() as u64;
409        let index_crc32c = crc32c::crc32c(&index_bytes);
410
411        self.inner.write_all(&index_bytes).await?;
412        self.current_offset += index_length;
413
414        // 40-byte trailer
415        let mut trailer = [0u8; 40];
416        trailer[0..8].copy_from_slice(&index_offset.to_le_bytes());
417        trailer[8..16].copy_from_slice(&index_length.to_le_bytes());
418        trailer[16..24].copy_from_slice(&kv_offset.to_le_bytes());
419        trailer[24..28].copy_from_slice(&kv_length.to_le_bytes());
420        trailer[28..32].copy_from_slice(&index_crc32c.to_le_bytes());
421        // [32..36] _reserved = 0
422        trailer[36..40].copy_from_slice(TRAILER_MAGIC);
423
424        self.inner.write_all(&trailer).await?;
425        self.inner.flush().await?;
426
427        Ok(self.inner)
428    }
429}
430
431// ── Encoding helpers ──────────────────────────────────────────────────────────
432
433fn validate_kv_keys(kv: &[(String, KvValue)]) -> Result<()> {
434    let mut seen = HashSet::new();
435    for (key, _) in kv {
436        if key.is_empty() {
437            return Err(Error::KvKeyEmpty);
438        }
439        if key.len() > 65535 {
440            return Err(Error::KvKeyTooLong { len: key.len() });
441        }
442        if !seen.insert(key.as_str()) {
443            return Err(Error::DuplicateKvKey(key.clone()));
444        }
445    }
446    Ok(())
447}
448
449fn encode_kv_section(kv: &[(String, KvValue)]) -> Result<Vec<u8>> {
450    let mut buf = Vec::new();
451    buf.extend_from_slice(&(kv.len() as u32).to_le_bytes());
452    for (key, value) in kv {
453        buf.extend_from_slice(&(key.len() as u16).to_le_bytes());
454        buf.extend_from_slice(key.as_bytes());
455        encode_kv_value(&mut buf, value)?;
456    }
457    Ok(buf)
458}
459
460fn kv_wire_tag(value: &KvValue) -> u8 {
461    match value {
462        KvValue::String(_) => 0x01,
463        KvValue::Int64(_) => 0x02,
464        KvValue::Uint64(_) => 0x03,
465        KvValue::Float64(_) => 0x04,
466        KvValue::Bool(_) => 0x05,
467        KvValue::Bytes(_) => 0x06,
468        KvValue::Array(_) => 0x07,
469    }
470}
471
472fn encode_kv_value(buf: &mut Vec<u8>, value: &KvValue) -> Result<()> {
473    buf.push(kv_wire_tag(value));
474    encode_kv_payload(buf, value)
475}
476
477fn encode_kv_payload(buf: &mut Vec<u8>, value: &KvValue) -> Result<()> {
478    match value {
479        KvValue::String(s) => {
480            buf.extend_from_slice(&(s.len() as u32).to_le_bytes());
481            buf.extend_from_slice(s.as_bytes());
482        }
483        KvValue::Int64(v) => buf.extend_from_slice(&v.to_le_bytes()),
484        KvValue::Uint64(v) => buf.extend_from_slice(&v.to_le_bytes()),
485        KvValue::Float64(v) => buf.extend_from_slice(&v.to_le_bytes()),
486        KvValue::Bool(v) => buf.push(u8::from(*v)),
487        KvValue::Bytes(v) => {
488            buf.extend_from_slice(&(v.len() as u32).to_le_bytes());
489            buf.extend_from_slice(v);
490        }
491        KvValue::Array(elements) => {
492            if elements.is_empty() {
493                return Err(Error::InvalidHeader("KV array must not be empty".into()));
494            }
495            let elem_tag = kv_wire_tag(&elements[0]);
496            if elem_tag == 0x07 {
497                return Err(Error::InvalidHeader(
498                    "KV array elements must not be arrays".into(),
499                ));
500            }
501            for elem in elements {
502                if kv_wire_tag(elem) != elem_tag {
503                    return Err(Error::InvalidHeader(
504                        "KV array elements must all have the same type".into(),
505                    ));
506                }
507            }
508            buf.push(elem_tag);
509            buf.extend_from_slice(&(elements.len() as u32).to_le_bytes());
510            for elem in elements {
511                encode_kv_payload(buf, elem)?;
512            }
513        }
514    }
515    Ok(())
516}
517
518fn encode_index_section(entries: &[InternalEntry]) -> Vec<u8> {
519    let mut buf = Vec::new();
520    buf.extend_from_slice(&(entries.len() as u64).to_le_bytes());
521    for e in entries {
522        buf.extend_from_slice(&(e.name.len() as u16).to_le_bytes());
523        buf.extend_from_slice(e.name.as_bytes());
524        buf.extend_from_slice(&e.descriptor_offset.to_le_bytes());
525        buf.extend_from_slice(&e.descriptor_length.to_le_bytes());
526        buf.extend_from_slice(&e.data_offset.to_le_bytes());
527        buf.extend_from_slice(&e.data_length.to_le_bytes());
528        buf.extend_from_slice(&0u32.to_le_bytes()); // flags, reserved
529    }
530    buf
531}