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
12pub const DEFAULT_MAX_COMPOSITE_DEPTH: usize = 64;
14
15#[derive(Debug)]
17pub struct FileTensor {
18 pub name: String,
20 pub descriptor: TensorDescriptor,
22 pub buffers: Vec<Bytes>,
24}
25
26#[derive(Debug)]
28pub enum FileItem {
29 Tensor(FileTensor),
31 Composite(FileComposite),
33}
34
35impl FileItem {
36 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#[derive(Debug)]
51pub struct FileComposite {
52 pub name: String,
54 pub head: TensorDescriptor,
56 pub members: Vec<FileItem>,
58}
59
60#[derive(Debug)]
82pub struct FileReader<R> {
83 inner: R,
84 data_buffer_alignment: u64,
85 max_composite_depth: usize,
86 pub index: Vec<IndexEntry>,
89 pub kv: Vec<(String, KvValue)>,
91}
92
93impl<R: AsyncRead + AsyncSeek + Unpin> FileReader<R> {
94 pub async fn open(mut inner: R) -> Result<Self> {
107 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 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 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 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 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 pub fn tensor_names(&self) -> impl Iterator<Item = &str> {
223 self.index.iter().map(|e| e.name.as_str())
224 }
225
226 pub fn kv(&self) -> &[(String, KvValue)] {
228 &self.kv
229 }
230
231 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 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 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 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 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 pub async fn read_composite(&mut self, head_name: &str) -> Result<FileComposite> {
332 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 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 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 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
416fn 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)?; 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}