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}