Skip to main content

hdf5_reader/
datatype_api.rs

1use crate::error::{ByteOrder, Error, Result};
2use crate::messages::datatype::Datatype;
3
4// Re-export types from the datatype message module so users don't need to
5// reach into messages::datatype.
6pub use crate::messages::datatype::{
7    CompoundField, EnumMember, ReferenceType, StringEncoding, StringPadding, StringSize, VarLenKind,
8};
9
10/// Trait for types that can be read from HDF5 datasets.
11///
12/// Implemented for primitive numeric types. Users can implement this
13/// for custom types (e.g., compound types).
14pub trait H5Type: Sized + Send + Clone {
15    /// The HDF5 datatype that this Rust type corresponds to.
16    fn hdf5_type() -> Datatype;
17
18    /// Decode a single value from raw bytes with the given datatype.
19    fn from_bytes(bytes: &[u8], dtype: &Datatype) -> Result<Self>;
20
21    /// Size of a single element in bytes.
22    fn element_size(dtype: &Datatype) -> usize;
23
24    /// Decode many values at once when the datatype has an efficient bulk path.
25    ///
26    /// Returning `None` falls back to per-element decoding.
27    fn decode_vec(_raw: &[u8], _dtype: &Datatype, _count: usize) -> Option<Result<Vec<Self>>> {
28        None
29    }
30
31    /// Whether raw bytes for this datatype can be copied directly into a `Vec<Self>`
32    /// without any further decoding or byte swapping.
33    fn native_copy_compatible(_dtype: &Datatype) -> bool {
34        false
35    }
36}
37
38/// Read a numeric value from bytes, handling byte-order conversion.
39fn read_numeric<const N: usize>(bytes: &[u8], byte_order: ByteOrder) -> Result<[u8; N]> {
40    if bytes.len() < N {
41        return Err(Error::InvalidData(format!(
42            "expected {} bytes, got {}",
43            N,
44            bytes.len()
45        )));
46    }
47    let mut arr = [0u8; N];
48    arr.copy_from_slice(&bytes[..N]);
49
50    // Swap bytes if the source endianness doesn't match native
51    #[cfg(target_endian = "little")]
52    if byte_order == ByteOrder::BigEndian {
53        arr.reverse();
54    }
55    #[cfg(target_endian = "big")]
56    if byte_order == ByteOrder::LittleEndian {
57        arr.reverse();
58    }
59
60    Ok(arr)
61}
62
63fn byte_order_is_native(byte_order: ByteOrder) -> bool {
64    #[cfg(target_endian = "little")]
65    {
66        byte_order == ByteOrder::LittleEndian
67    }
68    #[cfg(target_endian = "big")]
69    {
70        byte_order == ByteOrder::BigEndian
71    }
72}
73
74macro_rules! impl_h5type_int {
75    ($ty:ty, $size:literal, $signed:literal) => {
76        impl H5Type for $ty {
77            fn hdf5_type() -> Datatype {
78                Datatype::FixedPoint {
79                    size: $size,
80                    signed: $signed,
81                    byte_order: if cfg!(target_endian = "little") {
82                        ByteOrder::LittleEndian
83                    } else {
84                        ByteOrder::BigEndian
85                    },
86                }
87            }
88
89            fn from_bytes(bytes: &[u8], dtype: &Datatype) -> Result<Self> {
90                match dtype {
91                    Datatype::FixedPoint {
92                        size,
93                        signed,
94                        byte_order,
95                    } => {
96                        if *size as usize != std::mem::size_of::<$ty>() || *signed != $signed {
97                            return Err(Error::TypeMismatch {
98                                expected: stringify!($ty).into(),
99                                actual: format!("FixedPoint(size={}, signed={})", size, signed),
100                            });
101                        }
102                        let arr = read_numeric::<$size>(bytes, *byte_order)?;
103                        Ok(<$ty>::from_ne_bytes(arr))
104                    }
105                    _ => Err(Error::TypeMismatch {
106                        expected: stringify!($ty).into(),
107                        actual: format!("{:?}", dtype),
108                    }),
109                }
110            }
111
112            fn element_size(_dtype: &Datatype) -> usize {
113                $size
114            }
115
116            fn decode_vec(raw: &[u8], dtype: &Datatype, count: usize) -> Option<Result<Vec<Self>>> {
117                match dtype {
118                    Datatype::FixedPoint {
119                        size,
120                        signed,
121                        byte_order,
122                    } if *size as usize == $size && *signed == $signed => {
123                        let total_bytes = count.checked_mul($size)?;
124                        if raw.len() < total_bytes {
125                            return None;
126                        }
127
128                        let bytes = &raw[..total_bytes];
129                        if byte_order_is_native(*byte_order) {
130                            let mut values = Vec::<$ty>::with_capacity(count);
131                            // SAFETY: `values` has capacity for `count`
132                            // elements, `bytes` contains exactly
133                            // `count * size_of::<$ty>()` bytes, and integer
134                            // types admit every bit pattern.
135                            unsafe {
136                                std::ptr::copy_nonoverlapping(
137                                    bytes.as_ptr(),
138                                    values.as_mut_ptr() as *mut u8,
139                                    total_bytes,
140                                );
141                                values.set_len(count);
142                            }
143                            Some(Ok(values))
144                        } else {
145                            Some(Ok(bytes
146                                .chunks_exact($size)
147                                .map(|chunk| {
148                                    let mut arr = [0u8; $size];
149                                    arr.copy_from_slice(chunk);
150                                    arr.reverse();
151                                    <$ty>::from_ne_bytes(arr)
152                                })
153                                .collect()))
154                        }
155                    }
156                    _ => None,
157                }
158            }
159
160            fn native_copy_compatible(dtype: &Datatype) -> bool {
161                matches!(
162                    dtype,
163                    Datatype::FixedPoint {
164                        size,
165                        signed,
166                        byte_order,
167                    } if *size as usize == $size
168                        && *signed == $signed
169                        && byte_order_is_native(*byte_order)
170                )
171            }
172        }
173    };
174}
175
176impl_h5type_int!(i8, 1, true);
177impl_h5type_int!(u8, 1, false);
178impl_h5type_int!(i16, 2, true);
179impl_h5type_int!(u16, 2, false);
180impl_h5type_int!(i32, 4, true);
181impl_h5type_int!(u32, 4, false);
182impl_h5type_int!(i64, 8, true);
183impl_h5type_int!(u64, 8, false);
184
185impl H5Type for f32 {
186    fn hdf5_type() -> Datatype {
187        Datatype::FloatingPoint {
188            size: 4,
189            byte_order: if cfg!(target_endian = "little") {
190                ByteOrder::LittleEndian
191            } else {
192                ByteOrder::BigEndian
193            },
194        }
195    }
196
197    fn from_bytes(bytes: &[u8], dtype: &Datatype) -> Result<Self> {
198        match dtype {
199            Datatype::FloatingPoint { size, byte_order } => {
200                if *size != 4 {
201                    return Err(Error::TypeMismatch {
202                        expected: "f32".into(),
203                        actual: format!("FloatingPoint(size={})", size),
204                    });
205                }
206                let arr = read_numeric::<4>(bytes, *byte_order)?;
207                Ok(f32::from_ne_bytes(arr))
208            }
209            _ => Err(Error::TypeMismatch {
210                expected: "f32".into(),
211                actual: format!("{:?}", dtype),
212            }),
213        }
214    }
215
216    fn element_size(_dtype: &Datatype) -> usize {
217        4
218    }
219
220    fn decode_vec(raw: &[u8], dtype: &Datatype, count: usize) -> Option<Result<Vec<Self>>> {
221        match dtype {
222            Datatype::FloatingPoint { size, byte_order } if *size == 4 => {
223                let total_bytes = count.checked_mul(4)?;
224                if raw.len() < total_bytes {
225                    return None;
226                }
227
228                let bytes = &raw[..total_bytes];
229                if byte_order_is_native(*byte_order) {
230                    let mut values = Vec::<f32>::with_capacity(count);
231                    // SAFETY: `values` has capacity for `count` elements,
232                    // `bytes` contains exactly `count * size_of::<f32>()`
233                    // bytes, and `f32` admits every bit pattern.
234                    unsafe {
235                        std::ptr::copy_nonoverlapping(
236                            bytes.as_ptr(),
237                            values.as_mut_ptr() as *mut u8,
238                            total_bytes,
239                        );
240                        values.set_len(count);
241                    }
242                    Some(Ok(values))
243                } else {
244                    Some(Ok(bytes
245                        .chunks_exact(4)
246                        .map(|chunk| {
247                            let mut arr = [0u8; 4];
248                            arr.copy_from_slice(chunk);
249                            arr.reverse();
250                            f32::from_ne_bytes(arr)
251                        })
252                        .collect()))
253                }
254            }
255            _ => None,
256        }
257    }
258
259    fn native_copy_compatible(dtype: &Datatype) -> bool {
260        matches!(
261            dtype,
262            Datatype::FloatingPoint { size, byte_order }
263                if *size == 4 && byte_order_is_native(*byte_order)
264        )
265    }
266}
267
268impl H5Type for f64 {
269    fn hdf5_type() -> Datatype {
270        Datatype::FloatingPoint {
271            size: 8,
272            byte_order: if cfg!(target_endian = "little") {
273                ByteOrder::LittleEndian
274            } else {
275                ByteOrder::BigEndian
276            },
277        }
278    }
279
280    fn from_bytes(bytes: &[u8], dtype: &Datatype) -> Result<Self> {
281        match dtype {
282            Datatype::FloatingPoint { size, byte_order } => {
283                if *size != 8 {
284                    return Err(Error::TypeMismatch {
285                        expected: "f64".into(),
286                        actual: format!("FloatingPoint(size={})", size),
287                    });
288                }
289                let arr = read_numeric::<8>(bytes, *byte_order)?;
290                Ok(f64::from_ne_bytes(arr))
291            }
292            _ => Err(Error::TypeMismatch {
293                expected: "f64".into(),
294                actual: format!("{:?}", dtype),
295            }),
296        }
297    }
298
299    fn element_size(_dtype: &Datatype) -> usize {
300        8
301    }
302
303    fn decode_vec(raw: &[u8], dtype: &Datatype, count: usize) -> Option<Result<Vec<Self>>> {
304        match dtype {
305            Datatype::FloatingPoint { size, byte_order } if *size == 8 => {
306                let total_bytes = count.checked_mul(8)?;
307                if raw.len() < total_bytes {
308                    return None;
309                }
310
311                let bytes = &raw[..total_bytes];
312                if byte_order_is_native(*byte_order) {
313                    let mut values = Vec::<f64>::with_capacity(count);
314                    // SAFETY: `values` has capacity for `count` elements,
315                    // `bytes` contains exactly `count * size_of::<f64>()`
316                    // bytes, and `f64` admits every bit pattern.
317                    unsafe {
318                        std::ptr::copy_nonoverlapping(
319                            bytes.as_ptr(),
320                            values.as_mut_ptr() as *mut u8,
321                            total_bytes,
322                        );
323                        values.set_len(count);
324                    }
325                    Some(Ok(values))
326                } else {
327                    Some(Ok(bytes
328                        .chunks_exact(8)
329                        .map(|chunk| {
330                            let mut arr = [0u8; 8];
331                            arr.copy_from_slice(chunk);
332                            arr.reverse();
333                            f64::from_ne_bytes(arr)
334                        })
335                        .collect()))
336                }
337            }
338            _ => None,
339        }
340    }
341
342    fn native_copy_compatible(dtype: &Datatype) -> bool {
343        matches!(
344            dtype,
345            Datatype::FloatingPoint { size, byte_order }
346                if *size == 8 && byte_order_is_native(*byte_order)
347        )
348    }
349}
350
351/// Get the element size from a datatype.
352pub fn dtype_element_size(dtype: &Datatype) -> Result<usize> {
353    match dtype {
354        Datatype::FixedPoint { size, .. } => Ok(*size as usize),
355        Datatype::FloatingPoint { size, .. } => Ok(*size as usize),
356        Datatype::String {
357            size: StringSize::Fixed(n),
358            ..
359        } => Ok(*n as usize),
360        Datatype::String {
361            size: StringSize::Variable,
362            ..
363        } => Ok(16),
364        Datatype::Compound { size, .. } => Ok(*size as usize),
365        Datatype::Array { base, dims } => {
366            let base_size = dtype_element_size(base)?;
367            let count = dims.iter().try_fold(1usize, |acc, &dim| {
368                let dim = usize::try_from(dim).map_err(|_| {
369                    Error::InvalidData(
370                        "array datatype dimension exceeds platform usize capacity".to_string(),
371                    )
372                })?;
373                acc.checked_mul(dim).ok_or_else(|| {
374                    Error::InvalidData(
375                        "array datatype element count exceeds platform usize capacity".to_string(),
376                    )
377                })
378            })?;
379            base_size.checked_mul(count).ok_or_else(|| {
380                Error::InvalidData(
381                    "array datatype byte size exceeds platform usize capacity".to_string(),
382                )
383            })
384        }
385        Datatype::Enum { base, .. } => dtype_element_size(base),
386        Datatype::VarLen { .. } => Ok(16),
387        Datatype::Opaque { size, .. } => Ok(*size as usize),
388        Datatype::Reference { size, .. } => Ok(*size as usize),
389        Datatype::Bitfield { size, .. } => Ok(*size as usize),
390    }
391}
392
393#[cfg(test)]
394mod tests {
395    use super::*;
396
397    #[test]
398    fn f32_bulk_decode_native_endian() {
399        let dtype = <f32 as H5Type>::hdf5_type();
400        let raw = [0.5f32.to_ne_bytes(), 1.25f32.to_ne_bytes()].concat();
401        let values = <f32 as H5Type>::decode_vec(&raw, &dtype, 2)
402            .unwrap()
403            .unwrap();
404        assert_eq!(values, vec![0.5, 1.25]);
405    }
406
407    #[test]
408    fn u32_bulk_decode_big_endian() {
409        let dtype = Datatype::FixedPoint {
410            size: 4,
411            signed: false,
412            byte_order: ByteOrder::BigEndian,
413        };
414        let raw = [1u32.to_be_bytes(), 7u32.to_be_bytes()].concat();
415        let values = <u32 as H5Type>::decode_vec(&raw, &dtype, 2)
416            .unwrap()
417            .unwrap();
418        assert_eq!(values, vec![1, 7]);
419    }
420
421    #[test]
422    fn integer_from_bytes_rejects_signedness_mismatch() {
423        let dtype = Datatype::FixedPoint {
424            size: 2,
425            signed: false,
426            byte_order: ByteOrder::LittleEndian,
427        };
428
429        let err = <i16 as H5Type>::from_bytes(&u16::MAX.to_le_bytes(), &dtype).unwrap_err();
430        assert!(matches!(
431            err,
432            Error::TypeMismatch {
433                expected,
434                actual
435            } if expected == "i16" && actual.contains("signed=false")
436        ));
437    }
438
439    #[test]
440    fn integer_bulk_decode_rejects_signedness_mismatch() {
441        let unsigned_dtype = Datatype::FixedPoint {
442            size: 2,
443            signed: false,
444            byte_order: ByteOrder::LittleEndian,
445        };
446        let signed_dtype = Datatype::FixedPoint {
447            size: 2,
448            signed: true,
449            byte_order: ByteOrder::LittleEndian,
450        };
451
452        assert!(<i16 as H5Type>::decode_vec(&[0, 0], &unsigned_dtype, 1).is_none());
453        assert!(<u16 as H5Type>::decode_vec(&[0, 0], &signed_dtype, 1).is_none());
454    }
455
456    #[test]
457    fn integer_native_copy_compatible_rejects_signedness_mismatch() {
458        let unsigned_dtype = Datatype::FixedPoint {
459            size: 2,
460            signed: false,
461            byte_order: if cfg!(target_endian = "little") {
462                ByteOrder::LittleEndian
463            } else {
464                ByteOrder::BigEndian
465            },
466        };
467        let signed_dtype = Datatype::FixedPoint {
468            size: 2,
469            signed: true,
470            byte_order: if cfg!(target_endian = "little") {
471                ByteOrder::LittleEndian
472            } else {
473                ByteOrder::BigEndian
474            },
475        };
476
477        assert!(!<i16 as H5Type>::native_copy_compatible(&unsigned_dtype));
478        assert!(!<u16 as H5Type>::native_copy_compatible(&signed_dtype));
479        assert!(<u16 as H5Type>::native_copy_compatible(&unsigned_dtype));
480        assert!(<i16 as H5Type>::native_copy_compatible(&signed_dtype));
481    }
482
483    #[test]
484    fn dtype_element_size_rejects_array_overflow() {
485        let dtype = Datatype::Array {
486            base: Box::new(Datatype::FixedPoint {
487                size: 8,
488                signed: false,
489                byte_order: ByteOrder::LittleEndian,
490            }),
491            dims: vec![u64::MAX, 2],
492        };
493
494        let err = dtype_element_size(&dtype).unwrap_err();
495        assert!(err.to_string().contains("array datatype"));
496    }
497}