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
10pub enum FileCompositeNode<'a> {
14 Tensor {
16 name: &'a str,
18 descriptor: &'a TensorDescriptor,
20 buffers: &'a [&'a [u8]],
22 },
23 Composite {
25 name: &'a str,
27 head: &'a TensorDescriptor,
29 members: &'a [FileCompositeNode<'a>],
31 },
32}
33
34impl FileCompositeNode<'_> {
35 fn descriptor(&self) -> &TensorDescriptor {
37 match self {
38 FileCompositeNode::Tensor { descriptor, .. } => descriptor,
39 FileCompositeNode::Composite { head, .. } => head,
40 }
41 }
42}
43
44fn 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
78pub 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 pub async fn new(inner: W) -> Result<Self> {
159 Self::with_options(inner, FileWriterOptions::default()).await
160 }
161
162 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 let mut flags: u32 = file_flags::HAS_INDEX_CRC32C;
172 if sorted_index {
173 flags |= file_flags::SORTED_INDEX;
174 }
175
176 let mut header = [0u8; 64];
178 header[0..8].copy_from_slice(FILE_MAGIC);
179 header[8] = 1; header[9] = 0; header[12..16].copy_from_slice(&flags.to_le_bytes());
183 header[16..20].copy_from_slice(&(align as u32).to_le_bytes());
184 header[20..28].copy_from_slice(&FILE_HEADER_SIZE.to_le_bytes());
186 header[28..36].copy_from_slice(&u64::MAX.to_le_bytes());
188 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 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 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 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 let data_length = self.current_offset - data_offset;
275
276 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 pub async fn write_composite(
337 &mut self,
338 head_name: &str,
339 head: &TensorDescriptor,
340 members: &[FileCompositeNode<'_>],
341 ) -> Result<()> {
342 validate_composite_node(head, members)?;
344
345 self.write_tensor(head_name, head, &[]).await?;
347 self.write_composite_members(members).await
348 }
349
350 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::pin(self.write_composite_members(members)).await?;
369 }
370 }
371 }
372 Ok(())
373 }
374
375 pub async fn finish(mut self, kv: Vec<(String, KvValue)>) -> Result<W> {
380 validate_kv_keys(&kv)?;
381
382 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 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 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 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
431fn 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()); }
530 buf
531}