boreholeio/strided_array_file/
mod.rs1mod 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 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 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 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 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
206pub fn extract_array_index(uri: &ArrayUriReference) -> Result<Option<usize>, Error> {
214 if let Some(fragment) = uri.fragment() {
216 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#[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
287fn 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 std::slice::from_raw_parts(vec.as_ptr() as *const u8, std::mem::size_of_val(vec))
309 };
310 file.write_all(bytes)
311}