1use bytemuck::cast_slice;
2use memmap2::Mmap;
3use std::collections::HashMap;
4use std::fs;
5use std::path::Path;
6
7use crate::dtype::{DType, TrxScalar};
8use crate::error::{Result, TrxError};
9use crate::header::Header;
10use crate::io::filename::TrxFilename;
11use crate::mmap_backing::vec_to_bytes;
12use crate::mmap_backing::MmapBacking;
13use crate::trx_file::{DataArray, DataPerGroup, TrxFile, TrxParts};
14
15pub(crate) use super::zip::OffsetsDtype;
17
18fn mmap_file(path: &Path) -> Result<Mmap> {
20 let file = fs::File::open(path)?;
21 let mmap = unsafe { Mmap::map(&file)? };
23 Ok(mmap)
24}
25
26fn load_data_dir(dir: &Path) -> Result<HashMap<String, DataArray>> {
28 let mut map = HashMap::new();
29 if !dir.exists() {
30 return Ok(map);
31 }
32
33 for entry in fs::read_dir(dir)? {
34 let entry = entry?;
35 let path = entry.path();
36 if !path.is_file() {
37 continue;
38 }
39 let file_name = path
40 .file_name()
41 .and_then(|n| n.to_str())
42 .ok_or_else(|| TrxError::Format(format!("invalid filename: {}", path.display())))?;
43
44 let parsed = TrxFilename::parse(file_name)?;
45 let mmap = mmap_file(&path)?;
46
47 map.insert(
48 parsed.name.clone(),
49 DataArray::from_backing(MmapBacking::ReadOnly(mmap), parsed.ncols, parsed.dtype),
50 );
51 }
52
53 Ok(map)
54}
55
56fn load_dpg_dir(dir: &Path) -> Result<DataPerGroup> {
57 let mut out = HashMap::new();
58 if !dir.exists() {
59 return Ok(out);
60 }
61
62 for entry in fs::read_dir(dir)? {
63 let entry = entry?;
64 let path = entry.path();
65 if !path.is_dir() {
66 continue;
67 }
68 let group_name = entry.file_name().to_string_lossy().to_string();
69 let data = load_data_dir(&path)?;
70 if !data.is_empty() {
71 out.insert(group_name, data);
72 }
73 }
74
75 Ok(out)
76}
77
78fn find_file_with_prefix(dir: &Path, prefix: &str) -> Result<std::path::PathBuf> {
81 for entry in fs::read_dir(dir)? {
82 let entry = entry?;
83 let name = entry.file_name();
84 let name_str = name.to_string_lossy();
85 if name_str.starts_with(prefix) && name_str.chars().nth(prefix.len()) == Some('.') {
86 return Ok(entry.path());
87 }
88 }
89 Err(TrxError::FileNotFound(dir.join(prefix)))
90}
91
92pub fn load_from_directory<P: TrxScalar>(
94 dir: &Path,
95 tempdir: Option<tempfile::TempDir>,
96) -> Result<TrxFile<P>> {
97 if !dir.is_dir() {
98 return Err(TrxError::FileNotFound(dir.to_path_buf()));
99 }
100
101 let header = Header::from_file(&dir.join("header.json"))?;
103
104 let pos_path = find_file_with_prefix(dir, "positions")?;
106 let pos_fname = pos_path
107 .file_name()
108 .and_then(|n| n.to_str())
109 .ok_or_else(|| TrxError::Format("invalid positions filename".into()))?;
110 let pos_parsed = TrxFilename::parse(pos_fname)?;
111
112 if pos_parsed.dtype != P::DTYPE {
113 return Err(TrxError::DType(format!(
114 "expected positions dtype {}, got {}",
115 P::DTYPE,
116 pos_parsed.dtype
117 )));
118 }
119 if pos_parsed.ncols != 3 {
120 return Err(TrxError::Format(format!(
121 "positions must have 3 columns, got {}",
122 pos_parsed.ncols
123 )));
124 }
125
126 let positions_backing = MmapBacking::ReadOnly(mmap_file(&pos_path)?);
127
128 let off_path = find_file_with_prefix(dir, "offsets")?;
130 let off_fname = off_path
131 .file_name()
132 .and_then(|n| n.to_str())
133 .ok_or_else(|| TrxError::Format("invalid offsets filename".into()))?;
134 let off_parsed = TrxFilename::parse(off_fname)?;
135
136 let offsets_mmap = mmap_file(&off_path)?;
137 let offsets_backing = convert_offsets_to_u32(
138 &offsets_mmap,
139 off_parsed.dtype,
140 header.nb_streamlines as usize,
141 header.nb_vertices as usize,
142 )?;
143
144 let dps = load_data_dir(&dir.join("dps"))?;
146 let dpv = load_data_dir(&dir.join("dpv"))?;
147 let groups = load_data_dir(&dir.join("groups"))?;
148 let dpg = load_dpg_dir(&dir.join("dpg"))?;
149
150 Ok(TrxFile::from_parts(TrxParts {
151 header,
152 positions_backing,
153 offsets_backing,
154 dps,
155 dpv,
156 groups,
157 dpg,
158 tempdir,
159 }))
160}
161
162fn convert_offsets_to_u32(
165 mmap: &Mmap,
166 dtype: DType,
167 nb_streamlines: usize,
168 nb_vertices: usize,
169) -> Result<MmapBacking> {
170 match dtype {
171 DType::UInt64 => {
172 let values: &[u64] = cast_slice(mmap.as_ref());
173 if values.len() == nb_streamlines {
175 let mut owned: Vec<u32> = values
177 .iter()
178 .copied()
179 .map(|value| {
180 u32::try_from(value).map_err(|_| {
181 TrxError::Format(format!("offset {value} exceeds uint32 range"))
182 })
183 })
184 .collect::<Result<_>>()?;
185 owned.push(nb_vertices as u32);
186 let bytes: Vec<u8> = crate::mmap_backing::vec_to_bytes(owned);
187 Ok(MmapBacking::Owned(bytes))
188 } else if values.len() == nb_streamlines + 1 {
189 let owned: Vec<u32> = values
190 .iter()
191 .copied()
192 .map(|value| {
193 u32::try_from(value).map_err(|_| {
194 TrxError::Format(format!("offset {value} exceeds uint32 range"))
195 })
196 })
197 .collect::<Result<_>>()?;
198 Ok(MmapBacking::Owned(crate::mmap_backing::vec_to_bytes(owned)))
199 } else {
200 Err(TrxError::Format(format!(
201 "unexpected offset count: {} (expected {} or {})",
202 values.len(),
203 nb_streamlines,
204 nb_streamlines + 1,
205 )))
206 }
207 }
208 DType::UInt32 => {
209 let values: &[u32] = cast_slice(mmap.as_ref());
210 let mut out: Vec<u32> = values.to_vec();
211 if out.len() == nb_streamlines {
212 out.push(nb_vertices as u32);
213 }
214 let bytes: Vec<u8> = crate::mmap_backing::vec_to_bytes(out);
215 Ok(MmapBacking::Owned(bytes))
216 }
217 other => Err(TrxError::DType(format!(
218 "offsets must be uint32 or uint64, got {other}"
219 ))),
220 }
221}
222
223pub fn save_to_directory<P: TrxScalar>(trx: &TrxFile<P>, dir: &Path) -> Result<()> {
226 let offsets_dtype = OffsetsDtype::pick_for(trx.offsets());
227 fs::create_dir_all(dir)?;
228
229 trx.header().write_to(&dir.join("header.json"))?;
231
232 let pos_filename = format!("positions.3.{}", P::DTYPE.name());
234 fs::write(dir.join(&pos_filename), trx.positions_bytes())?;
235
236 let offsets_filename = format!("offsets.{}", offsets_dtype.suffix());
238 let offsets_bytes = offsets_dtype.encode(trx.offsets());
239 fs::write(dir.join(offsets_filename), offsets_bytes)?;
240
241 save_data_dir(trx.dps_arrays(), &dir.join("dps"))?;
243
244 save_data_dir(trx.dpv_arrays(), &dir.join("dpv"))?;
246
247 save_data_dir(trx.group_arrays(), &dir.join("groups"))?;
249
250 save_dpg_dir(trx.dpg_arrays(), &dir.join("dpg"))?;
252
253 Ok(())
254}
255
256pub fn append_dps_to_directory(
258 dir: &Path,
259 dps: &HashMap<String, DataArray>,
260 overwrite: bool,
261) -> Result<()> {
262 let header = Header::from_file(&dir.join("header.json"))?;
263 validate_row_count("DPS", dps, header.nb_streamlines as usize)?;
264 append_arrays_to_directory(&dir.join("dps"), dps, overwrite)
265}
266
267pub fn append_dpv_to_directory(
269 dir: &Path,
270 dpv: &HashMap<String, DataArray>,
271 overwrite: bool,
272) -> Result<()> {
273 let header = Header::from_file(&dir.join("header.json"))?;
274 validate_row_count("DPV", dpv, header.nb_vertices as usize)?;
275 append_arrays_to_directory(&dir.join("dpv"), dpv, overwrite)
276}
277
278pub fn append_groups_to_directory(
280 dir: &Path,
281 groups: &HashMap<String, Vec<u32>>,
282 overwrite: bool,
283) -> Result<()> {
284 let header = Header::from_file(&dir.join("header.json"))?;
285 let groups_dir = dir.join("groups");
286 fs::create_dir_all(&groups_dir)?;
287 for (name, members) in groups {
288 validate_group_members(name, members, header.nb_streamlines as usize)?;
289 let target = groups_dir.join(format!("{name}.uint32"));
290 if !overwrite {
291 if let Some(existing) = find_named_array_file(&groups_dir, name)? {
292 if existing.exists() {
293 continue;
294 }
295 }
296 } else if let Some(existing) = find_named_array_file(&groups_dir, name)? {
297 if existing != target && existing.exists() {
298 fs::remove_file(existing)?;
299 }
300 }
301 fs::write(target, vec_to_bytes(members.clone()))?;
302 }
303 Ok(())
304}
305
306pub fn append_dpg_to_directory(dir: &Path, dpg: &DataPerGroup, overwrite: bool) -> Result<()> {
308 let groups_dir = dir.join("groups");
309 let dpg_root = dir.join("dpg");
310 for (group, entries) in dpg {
311 if find_named_array_file(&groups_dir, group)?.is_none() {
312 return Err(TrxError::Argument(format!(
313 "cannot add DPG entries for missing group '{group}'"
314 )));
315 }
316 let group_dir = dpg_root.join(group);
317 fs::create_dir_all(&group_dir)?;
318 for (name, arr) in entries {
319 let target = group_dir.join(filename_for_array(name, arr));
320 if !overwrite {
321 if let Some(existing) = find_named_array_file(&group_dir, name)? {
322 if existing.exists() {
323 continue;
324 }
325 }
326 } else if let Some(existing) = find_named_array_file(&group_dir, name)? {
327 if existing != target && existing.exists() {
328 fs::remove_file(existing)?;
329 }
330 }
331 fs::write(target, arr.as_bytes())?;
332 }
333 }
334 Ok(())
335}
336
337pub fn delete_dps_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
339 delete_named_arrays(&dir.join("dps"), names)
340}
341
342pub fn delete_dpv_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
344 delete_named_arrays(&dir.join("dpv"), names)
345}
346
347pub fn delete_groups_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
349 let groups_dir = dir.join("groups");
350 for name in names {
351 if let Some(path) = find_named_array_file(&groups_dir, name)? {
352 if path.exists() {
353 fs::remove_file(path)?;
354 }
355 }
356 let dpg_group = dir.join("dpg").join(name);
357 if dpg_group.exists() {
358 fs::remove_dir_all(dpg_group)?;
359 }
360 }
361 Ok(())
362}
363
364pub fn delete_dpg_from_directory(dir: &Path, group: &str, names: Option<&[&str]>) -> Result<()> {
369 let group_dir = dir.join("dpg").join(group);
370 match names {
371 None | Some([]) => {
372 if group_dir.exists() {
373 fs::remove_dir_all(group_dir)?;
374 }
375 }
376 Some(names) => {
377 for name in names {
378 if let Some(path) = find_named_array_file(&group_dir, name)? {
379 if path.exists() {
380 fs::remove_file(path)?;
381 }
382 }
383 }
384 }
385 }
386 Ok(())
387}
388
389fn save_data_dir(arrays: &HashMap<String, DataArray>, dir: &Path) -> Result<()> {
390 if arrays.is_empty() {
391 return Ok(());
392 }
393 fs::create_dir_all(dir)?;
394 for (name, arr) in arrays {
395 let filename = filename_for_array(name, arr);
396 fs::write(dir.join(&filename), arr.as_bytes())?;
397 }
398 Ok(())
399}
400
401fn save_dpg_dir(arrays: &DataPerGroup, dir: &Path) -> Result<()> {
402 if arrays.is_empty() {
403 return Ok(());
404 }
405 fs::create_dir_all(dir)?;
406 for (group, entries) in arrays {
407 save_data_dir(entries, &dir.join(group))?;
408 }
409 Ok(())
410}
411
412fn append_arrays_to_directory(
413 dir: &Path,
414 arrays: &HashMap<String, DataArray>,
415 overwrite: bool,
416) -> Result<()> {
417 fs::create_dir_all(dir)?;
418 for (name, arr) in arrays {
419 let target = dir.join(filename_for_array(name, arr));
420 if !overwrite {
421 if let Some(existing) = find_named_array_file(dir, name)? {
422 if existing.exists() {
423 continue;
424 }
425 }
426 } else if let Some(existing) = find_named_array_file(dir, name)? {
427 if existing != target && existing.exists() {
428 fs::remove_file(existing)?;
429 }
430 }
431 fs::write(target, arr.as_bytes())?;
432 }
433 Ok(())
434}
435
436fn delete_named_arrays(dir: &Path, names: &[&str]) -> Result<()> {
437 for name in names {
438 if let Some(path) = find_named_array_file(dir, name)? {
439 if path.exists() {
440 fs::remove_file(path)?;
441 }
442 }
443 }
444 Ok(())
445}
446
447fn find_named_array_file(dir: &Path, name: &str) -> Result<Option<std::path::PathBuf>> {
448 if !dir.exists() {
449 return Ok(None);
450 }
451 for entry in fs::read_dir(dir)? {
452 let entry = entry?;
453 let path = entry.path();
454 if !path.is_file() {
455 continue;
456 }
457 let file_name = path
458 .file_name()
459 .and_then(|n| n.to_str())
460 .ok_or_else(|| TrxError::Format(format!("invalid filename: {}", path.display())))?;
461 let parsed = TrxFilename::parse(file_name)?;
462 if parsed.name == name {
463 return Ok(Some(path));
464 }
465 }
466 Ok(None)
467}
468
469fn validate_row_count(
470 kind: &str,
471 arrays: &HashMap<String, DataArray>,
472 expected_rows: usize,
473) -> Result<()> {
474 for (name, arr) in arrays {
475 if arr.nrows() != expected_rows {
476 return Err(TrxError::Format(format!(
477 "{kind} '{name}' has {} rows, expected {expected_rows}",
478 arr.nrows()
479 )));
480 }
481 }
482 Ok(())
483}
484
485fn validate_group_members(name: &str, members: &[u32], nb_streamlines: usize) -> Result<()> {
486 for &member in members {
487 if member as usize >= nb_streamlines {
488 return Err(TrxError::Format(format!(
489 "group '{name}' contains streamline index {member}, but NB_STREAMLINES is {nb_streamlines}"
490 )));
491 }
492 }
493 Ok(())
494}
495
496fn filename_for_array(name: &str, arr: &DataArray) -> String {
497 TrxFilename {
498 name: name.to_string(),
499 ncols: arr.ncols(),
500 dtype: arr.dtype(),
501 }
502 .to_filename()
503}