Skip to main content

hurray_io/stream/
writer.rs

1use hurray_core::{CompositeValidator, SyncMode, TensorDescriptor};
2use tokio::io::{AsyncWrite, AsyncWriteExt};
3
4use crate::{Error, Result};
5
6/// A node to emit as part of a composite via [`StreamWriter::write_composite`]: either a
7/// plain tensor (its descriptor plus one byte-slice per buffer) or a nested composite (its
8/// head plus its own members). The recursive [`Composite`][CompositeNode::Composite] arm
9/// lets a member be a composite in its own right (ADR-027 § Binding).
10pub enum CompositeNode<'a> {
11    /// A single (non-composite) tensor: descriptor + one byte-slice per declared buffer.
12    Tensor {
13        /// The member's tensor descriptor.
14        descriptor: &'a TensorDescriptor,
15        /// One byte-slice per buffer handle in `descriptor.buffers`.
16        buffers: &'a [&'a [u8]],
17    },
18    /// A nested composite: a data-less head plus its ordered members.
19    Composite {
20        /// The nested composite's head descriptor.
21        head: &'a TensorDescriptor,
22        /// The nested composite's members, in order.
23        members: &'a [CompositeNode<'a>],
24    },
25}
26
27impl CompositeNode<'_> {
28    /// The node's governing descriptor: the tensor's descriptor, or the nested head.
29    fn descriptor(&self) -> &TensorDescriptor {
30        match self {
31            CompositeNode::Tensor { descriptor, .. } => descriptor,
32            CompositeNode::Composite { head, .. } => head,
33        }
34    }
35}
36
37/// Writes tensors to an async sink in the Hurray streaming wire format.
38///
39/// The wire format is bare concatenation: each tensor is represented as its
40/// encoded [`TensorDescriptor`] immediately followed by each buffer's raw bytes.
41/// No outer framing, no alignment padding between tensors.
42///
43/// # Examples
44///
45/// ```no_run
46/// # #[tokio::main]
47/// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
48/// use hurray_core::{
49///     BufferHandle, DeviceTag, ElementType, LayoutDescriptor, Shape, SyncMode,
50///     TensorDescriptor, MIN_BUFFER_ALIGNMENT,
51/// };
52/// use hurray_io::stream::StreamWriter;
53///
54/// let handle = BufferHandle::new(64, MIN_BUFFER_ALIGNMENT, DeviceTag::Cpu, SyncMode::ProducerSynced)?;
55/// let shape = Shape::new(vec![4u64, 4]).unwrap();
56/// let desc = TensorDescriptor::new(
57///     1, 0,
58///     ElementType::Float32,
59///     shape,
60///     0,
61///     LayoutDescriptor::RowMajor,
62///     vec![handle],
63///     None, None, None, None,
64/// )?;
65/// let data = vec![0u8; 64];
66///
67/// let mut wire = Vec::<u8>::new();
68/// let mut writer = StreamWriter::new(&mut wire);
69/// writer.write_tensor(&desc, &[&data]).await?;
70/// writer.finish().await?;
71/// # Ok(())
72/// # }
73/// ```
74pub struct StreamWriter<W> {
75    inner: W,
76    enforce_cross_machine_sync: bool,
77}
78
79impl<W: AsyncWrite + Unpin> StreamWriter<W> {
80    /// Creates a writer that accepts any valid [`SyncMode`].
81    pub fn new(inner: W) -> Self {
82        Self {
83            inner,
84            enforce_cross_machine_sync: false,
85        }
86    }
87
88    /// Creates a writer that rejects buffers whose `sync_mode` is not
89    /// [`SyncMode::ProducerSynced`].
90    ///
91    /// Use this when the stream crosses machine boundaries where GPU and
92    /// semaphore-based sync primitives are not transferable.
93    pub fn cross_machine(inner: W) -> Self {
94        Self {
95            inner,
96            enforce_cross_machine_sync: true,
97        }
98    }
99
100    /// Encodes and writes one tensor.
101    ///
102    /// # Errors
103    ///
104    /// - [`Error::MultiBufferLengthMismatch`] — `buffers.len()` ≠ `desc.buffers.len()`
105    /// - [`Error::BufferSizeMismatch`] — a buffer's byte length ≠ its handle's `byte_size`
106    /// - [`Error::InvalidCrossMachineSyncMode`] — cross-machine mode and a buffer has a
107    ///   non-`ProducerSynced` sync mode
108    /// - [`Error::Core`] — descriptor encoding failed
109    /// - [`Error::Io`] — underlying write error
110    pub async fn write_tensor(&mut self, desc: &TensorDescriptor, buffers: &[&[u8]]) -> Result<()> {
111        if buffers.len() != desc.buffers.len() {
112            return Err(Error::MultiBufferLengthMismatch {
113                declared: desc.buffers.len(),
114                actual: buffers.len(),
115            });
116        }
117
118        for (i, (handle, buf)) in desc.buffers.iter().zip(buffers).enumerate() {
119            let declared = handle.byte_size();
120            let actual = buf.len() as u64;
121            if declared != actual {
122                return Err(Error::BufferSizeMismatch {
123                    index: i,
124                    declared,
125                    actual,
126                });
127            }
128            if self.enforce_cross_machine_sync && handle.sync_mode() != SyncMode::ProducerSynced {
129                return Err(Error::InvalidCrossMachineSyncMode {
130                    index: i,
131                    actual: handle.sync_mode().to_byte(),
132                });
133            }
134        }
135
136        let encoded = desc.encode()?;
137        self.inner.write_all(&encoded).await?;
138
139        for buf in buffers {
140            self.inner.write_all(buf).await?;
141        }
142
143        Ok(())
144    }
145
146    /// Encodes and writes a composite tensor: its head followed by every member's
147    /// descriptor and data, in order (ADR-027 § Binding).
148    ///
149    /// The group is validated *before* any byte is written — reusing
150    /// [`CompositeValidator`] to check `member_count` and the per-rule constraints
151    /// (partition exact-cover / non-overlap, overlay base-first ordering) — so a torn or
152    /// invalid composite never reaches the wire. Members that are themselves composites are
153    /// written recursively (head precedes its members at every level), preserving the
154    /// forward, self-delimiting, back-reference-free wire contract.
155    ///
156    /// # Examples
157    ///
158    /// ```no_run
159    /// # #[tokio::main]
160    /// # async fn main() -> Result<(), Box<dyn std::error::Error>> {
161    /// use hurray_core::{
162    ///     layout::{CompositeLayout, CompositionRule, LayoutDescriptor},
163    ///     ElementType, Shape, ShardDescriptor, TensorDescriptor,
164    /// };
165    /// use hurray_io::stream::{CompositeNode, StreamWriter};
166    ///
167    /// // A partition head [8, 8] split into two [8, 4] members (buffers elided here).
168    /// let head = TensorDescriptor::new(
169    ///     1, 0, ElementType::Float32, Shape::new(vec![8u64, 8]).unwrap(), 0,
170    ///     LayoutDescriptor::Composite(CompositeLayout::new(CompositionRule::Partition, 2).unwrap()),
171    ///     vec![], None, None, None, None,
172    /// )?;
173    /// # let member_descs: Vec<TensorDescriptor> = vec![];
174    /// # let members: Vec<CompositeNode> = vec![];
175    /// let mut wire = Vec::<u8>::new();
176    /// let mut writer = StreamWriter::new(&mut wire);
177    /// writer.write_composite(&head, &members).await?;
178    /// writer.finish().await?;
179    /// # Ok(())
180    /// # }
181    /// ```
182    ///
183    /// # Errors
184    ///
185    /// - [`Error::Core`] — the head is not a valid composite head, or validation failed
186    ///   (member-count mismatch, partition coverage, overlay ordering)
187    /// - the same buffer/sync errors as [`write_tensor`][StreamWriter::write_tensor] for
188    ///   each member
189    /// - [`Error::Io`] — underlying write error
190    pub async fn write_composite(
191        &mut self,
192        head: &TensorDescriptor,
193        members: &[CompositeNode<'_>],
194    ) -> Result<()> {
195        // Validate every level up front, so a torn/invalid composite — at any nesting
196        // depth — never reaches the wire.
197        validate_composite_node(head, members)?;
198        self.write_composite_unchecked(head, members).await
199    }
200
201    /// Writes an already-validated composite (head, then each member recursively).
202    async fn write_composite_unchecked(
203        &mut self,
204        head: &TensorDescriptor,
205        members: &[CompositeNode<'_>],
206    ) -> Result<()> {
207        // The head owns no data buffers (composite-head invariant).
208        self.write_tensor(head, &[]).await?;
209
210        for member in members {
211            match member {
212                CompositeNode::Tensor {
213                    descriptor,
214                    buffers,
215                } => {
216                    self.write_tensor(descriptor, buffers).await?;
217                }
218                CompositeNode::Composite {
219                    head: nested_head,
220                    members: nested_members,
221                } => {
222                    // Already validated by write_composite's up-front pass; box the
223                    // recursive call since an async fn cannot name its own future.
224                    Box::pin(self.write_composite_unchecked(nested_head, nested_members)).await?;
225                }
226            }
227        }
228
229        Ok(())
230    }
231
232    /// Flushes the underlying sink and returns it.
233    pub async fn finish(mut self) -> Result<W> {
234        self.inner.flush().await?;
235        Ok(self.inner)
236    }
237}
238
239/// Recursively validates a composite node and all nested composites, writing nothing.
240///
241/// Reuses [`CompositeValidator`] at each level (member count + partition coverage / overlay
242/// ordering). Pure and synchronous, so [`StreamWriter::write_composite`] can validate the
243/// entire tree before emitting a single byte.
244fn validate_composite_node(head: &TensorDescriptor, members: &[CompositeNode<'_>]) -> Result<()> {
245    let mut validator = CompositeValidator::new(head)?;
246    for member in members {
247        validator.push_member(member.descriptor())?;
248    }
249    validator.finish()?;
250
251    for member in members {
252        if let CompositeNode::Composite {
253            head: nested_head,
254            members: nested_members,
255        } = member
256        {
257            validate_composite_node(nested_head, nested_members)?;
258        }
259    }
260    Ok(())
261}