Skip to main content

boreholeio/strided_array_file/
mod.rs

1mod array_proxy;
2mod typed_array;
3mod utils;
4
5use crate::schema::array_uri_reference::ArrayUriReference;
6use ndarray::ArrayViewD;
7use std::{
8    io::{Read, Write},
9    path::{Path, PathBuf},
10    sync::RwLock,
11};
12
13pub(crate) use utils::{convert, convert_vec};
14
15pub use array_proxy::{ArrayProxy, ProxyCreator, ProxyFunctor, Statistics};
16pub use typed_array::TypedArray;
17pub use utils::Error;
18
19#[derive(Debug)]
20pub struct StridedArrayFile<'a> {
21    path: PathBuf,
22    file_mmap: RwLock<Option<memmap2::Mmap>>,
23    arrays: RwLock<Option<Vec<ArrayProxy<'a>>>>,
24}
25
26impl<'a> StridedArrayFile<'a> {
27    pub fn new(path: PathBuf) -> Self {
28        Self {
29            path,
30            file_mmap: RwLock::new(None),
31            arrays: RwLock::new(None),
32        }
33    }
34
35    pub fn is_open(&self) -> bool {
36        self.file_mmap.read().unwrap().is_some() && self.arrays.read().unwrap().is_some()
37    }
38
39    pub fn write(file_path: &Path, arrays: &[ArrayProxy<'a>]) -> Result<(), std::io::Error> {
40        Self::validate_system_endianness()?;
41        let mut f = std::fs::OpenOptions::new()
42            .write(true)
43            .create(true)
44            .truncate(true)
45            .open(file_path)?;
46        Header {
47            n_arrays: convert(arrays.len()),
48        }
49        .write(&mut f)?;
50        let header_size = 10
51            + arrays
52                .iter()
53                .map(|arr| 13 + 16 * arr.dimensions())
54                .sum::<usize>();
55        struct ArrayDataParams {
56            start: usize,
57            size: usize,
58        }
59        let mut array_data_params = Vec::<ArrayDataParams>::with_capacity(arrays.len());
60        let mut data_start = header_size;
61        for a in arrays {
62            f.write_all(&convert::<usize, u32>(a.dimensions()).to_le_bytes())?;
63            let shape = a
64                .shape()
65                .iter()
66                .map(|s| convert::<usize, u64>(*s))
67                .collect::<Vec<u64>>();
68            write_vec(&shape, &mut f)?;
69            let elem_size = a.element_size();
70            data_start = data_start.div_ceil(elem_size) * elem_size;
71            let strides_in_bytes = convert_vec::<isize, i64>(&a.strides(true));
72            write_vec(&strides_in_bytes, &mut f)?;
73            f.write_all(&a.element_type().to_le_bytes())?;
74            // TODO take negative strides (span/data_start) into account
75            f.write_all(&convert::<usize, u64>(data_start).to_le_bytes())?;
76            let data_size = a.data_size();
77            array_data_params.push(ArrayDataParams {
78                start: data_start,
79                size: data_size,
80            });
81            data_start += data_size;
82        }
83        let mut cursor = header_size;
84        for (arr, params) in std::iter::zip(arrays, array_data_params) {
85            if params.start > cursor {
86                f.write_all(&vec![0u8; params.start - cursor][..])?;
87            }
88            f.write_all(unsafe {
89                // Note that the from_raw_parts method requires the memory to be nicely
90                // aligned, which is guaranteed in this particular case
91                std::slice::from_raw_parts(arr.data(), params.size)
92            })?;
93            cursor = params.start + params.size;
94        }
95        Ok(())
96    }
97
98    pub fn path(&self) -> &Path {
99        &self.path
100    }
101
102    pub fn count(&self) -> Result<usize, Error> {
103        if !self.is_open() {
104            self.open()?
105        }
106        Ok(self.arrays.read().unwrap().as_ref().unwrap().len())
107    }
108
109    pub fn try_get(&self, i: usize) -> Result<ArrayProxy<'a>, Error> {
110        if i >= self.count()? {
111            Err(format!(
112                "Array index {i} is greater than file size {}",
113                self.count()?
114            ))
115        } else {
116            Ok(self.arrays.read().unwrap().as_ref().unwrap()[i].clone())
117        }
118    }
119
120    pub fn try_get_as<T: 'static>(&self, i: usize) -> Result<ArrayViewD<'a, T>, Error> {
121        self.try_get(i)?.try_as::<T>()
122    }
123
124    pub fn get(&self, i: usize) -> ArrayProxy<'a> {
125        self.try_get(i).unwrap()
126    }
127
128    pub fn get_as<T: 'static>(&self, i: usize) -> ArrayViewD<'a, T> {
129        self.try_get_as(i).unwrap()
130    }
131
132    fn validate_system_endianness() -> Result<(), std::io::Error> {
133        // https://doc.rust-lang.org/reference/conditional-compilation.html#target_endian
134        if !cfg!(target_endian = "little") {
135            return Err(std::io::Error::new(
136                std::io::ErrorKind::Unsupported,
137                "The OS is not little endian",
138            ));
139        }
140        Ok(())
141    }
142
143    fn open(&self) -> Result<(), Error> {
144        if let Err(err) = self.open_unchecked() {
145            self.arrays.write().unwrap().take();
146            self.file_mmap.write().unwrap().take();
147            Err(err)
148        } else {
149            Ok(())
150        }
151    }
152
153    /// This method does not clean up `self` if an error occurs midway. One
154    /// must always call `self.open` instead.
155    fn open_unchecked(&self) -> Result<(), Error> {
156        Self::validate_system_endianness().map_err(|e| e.to_string())?;
157        if !self.path.is_file() {
158            return Err(format!(
159                "Invalid path to a Strided Array File: {:?}",
160                self.path
161            ));
162        }
163        let mut f = std::fs::OpenOptions::new()
164            .read(true)
165            .create(false)
166            .open(&self.path)
167            .map_err(|e| e.to_string())?;
168        let mmap = unsafe { memmap2::Mmap::map(&f).unwrap() };
169        let n_arrays = Header::read(&mut f).map_err(|e| e.to_string())?.n_arrays;
170        let mut offset = Header::N_BYTES;
171        let mut arrays = Vec::<ArrayProxy>::new();
172        for _ in 0..n_arrays {
173            let n_dims = convert(u32::from_le_bytes(
174                mmap[offset..offset + 4].try_into().unwrap(),
175            ));
176            offset += 4;
177            let shape = unpack_vec_non_aligned::<u64, _>(
178                &mmap[offset..offset + 8 * n_dims],
179                n_dims,
180                |bytes| u64::from_le_bytes(bytes.try_into().unwrap()),
181            );
182            offset += 8 * n_dims;
183            let strides_in_bytes = unpack_vec_non_aligned::<i64, _>(
184                &mmap[offset..offset + 8 * n_dims],
185                n_dims,
186                |bytes| i64::from_le_bytes(bytes.try_into().unwrap()),
187            );
188            offset += 8 * n_dims;
189            let element_type = mmap[offset];
190            let data_start = u64::from_le_bytes(mmap[offset + 1..offset + 9].try_into().unwrap());
191            offset += 9;
192            arrays.push(ArrayProxy::from_mmap(
193                &mmap,
194                shape,
195                strides_in_bytes,
196                element_type,
197                data_start,
198            )?);
199        }
200        self.file_mmap.write().unwrap().replace(mmap);
201        self.arrays.write().unwrap().replace(arrays);
202        Ok(())
203    }
204}
205
206/// Returns the array index (if present) specified in the URI.
207///
208/// For example:
209/// - Returns `Ok(None)` for `"some/path.star"`,
210/// - Returns `Ok(None)` for `"some/path.png"`,
211/// - Returns `Ok(Some(10))` for `"https://example.com/some/path.star#10"`,
212/// - Returns `Err` for `"some/path.png#eleven"`.
213pub fn extract_array_index(uri: &ArrayUriReference) -> Result<Option<usize>, Error> {
214    // See if there is anything after `#`.
215    if let Some(fragment) = uri.fragment() {
216        // Convert `&str` to `usize`.
217        Ok(Some(fragment.as_str().parse::<usize>().map_err(|e| {
218            format!("{fragment} is not a valid integer: {e}")
219        })?))
220    } else {
221        Ok(None)
222    }
223}
224
225#[derive(Debug, PartialEq)]
226struct Version {
227    major: u8,
228    minor: u8,
229}
230
231// C-represented header will introduce some padding. Adding macroses:
232// #[repr(packed(1))]
233// #[derive(Debug, serde::Deserialize)]
234// generates an error:
235// cannot move out of `self.version` which is behind a shared reference
236//`#[derive(Debug)]` triggers a move because taking references to the fields of a packed
237// struct is undefined behaviour.
238#[derive(Debug)]
239struct Header {
240    n_arrays: u32,
241}
242
243impl Header {
244    const MAGIC_BYTES: &'static [u8; 4] = b"StAr";
245    const VERSION: Version = Version { major: 0, minor: 2 };
246    const N_BYTES: usize = 10;
247
248    fn read(f: &mut std::fs::File) -> Result<Self, std::io::Error> {
249        assert_eq!(Self::N_BYTES, 10);
250        let mut magic_bytes = [0; 4];
251        f.read_exact(&mut magic_bytes)?;
252        if &magic_bytes != Self::MAGIC_BYTES {
253            return Err(std::io::Error::new(
254                std::io::ErrorKind::InvalidData,
255                "Magic bytes validation failed",
256            ));
257        }
258        let mut version = [0; 2];
259        f.read_exact(&mut version)?;
260        let version = Version {
261            major: version[0],
262            minor: version[1],
263        };
264        if version != Self::VERSION {
265            return Err(std::io::Error::new(
266                std::io::ErrorKind::InvalidData,
267                format!("Unsupported version: {version:?}"),
268            ));
269        }
270        let mut n_arrays = [0; 4];
271        f.read_exact(&mut n_arrays)?;
272        Ok(Self {
273            n_arrays: u32::from_le_bytes(n_arrays),
274        })
275    }
276
277    fn write(&self, f: &mut std::fs::File) -> Result<(), std::io::Error> {
278        assert_eq!(Self::N_BYTES, 10);
279        f.write_all(Self::MAGIC_BYTES)?;
280        f.write_all(&Self::VERSION.major.to_le_bytes())?;
281        f.write_all(&Self::VERSION.minor.to_le_bytes())?;
282        f.write_all(&self.n_arrays.to_le_bytes())?;
283        Ok(())
284    }
285}
286
287// This method is used only twice: with converter being either `i64::from_le_bytes` or
288// `u64::from_le_bytes`. The thing is that, those 2 methods don't have any common trait
289// and, thus, needs to be passed as a closure.
290fn unpack_vec_non_aligned<T, F>(bytes: &[u8], n_elems: usize, converter: F) -> Vec<T>
291where
292    F: Fn(&[u8]) -> T,
293{
294    let elem_size = std::mem::size_of::<T>();
295    let mut vec = Vec::<T>::with_capacity(n_elems);
296    let mut offset = 0;
297    for _ in 0..n_elems {
298        vec.push(converter(&bytes[offset..offset + elem_size]));
299        offset += elem_size;
300    }
301    vec
302}
303
304fn write_vec<T>(vec: &[T], file: &mut std::fs::File) -> Result<(), std::io::Error> {
305    let bytes: &[u8] = unsafe {
306        // Note that the from_raw_parts method requires the memory to be nicely
307        // aligned, which is guaranteed in this particular case
308        std::slice::from_raw_parts(vec.as_ptr() as *const u8, std::mem::size_of_val(vec))
309    };
310    file.write_all(bytes)
311}