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>(dir: &Path) -> Result<TrxFile<P>> {
94 if !dir.is_dir() {
95 return Err(TrxError::FileNotFound(dir.to_path_buf()));
96 }
97
98 let header = Header::from_file(&dir.join("header.json"))?;
100
101 let positions_backing = match find_file_with_prefix(dir, "positions") {
102 Ok(pos_path) => {
103 let pos_fname = pos_path
104 .file_name()
105 .and_then(|n| n.to_str())
106 .ok_or_else(|| TrxError::Format("invalid positions filename".into()))?;
107 let pos_parsed = TrxFilename::parse(pos_fname)?;
108
109 if pos_parsed.dtype != P::DTYPE {
110 return Err(TrxError::DType(format!(
111 "expected positions dtype {}, got {}",
112 P::DTYPE,
113 pos_parsed.dtype
114 )));
115 }
116 if pos_parsed.ncols != 3 {
117 return Err(TrxError::Format(format!(
118 "positions must have 3 columns, got {}",
119 pos_parsed.ncols
120 )));
121 }
122 MmapBacking::ReadOnly(mmap_file(&pos_path)?)
123 }
124 Err(e) => {
125 if header.nb_vertices == 0 {
126 MmapBacking::Owned(Vec::new())
127 } else {
128 return Err(e);
129 }
130 }
131 };
132
133 let offsets_backing = match find_file_with_prefix(dir, "offsets") {
134 Ok(off_path) => {
135 let off_fname = off_path
136 .file_name()
137 .and_then(|n| n.to_str())
138 .ok_or_else(|| TrxError::Format("invalid offsets filename".into()))?;
139 let off_parsed = TrxFilename::parse(off_fname)?;
140
141 let offsets_mmap = mmap_file(&off_path)?;
142 convert_offsets_to_u32(
143 &offsets_mmap,
144 off_parsed.dtype,
145 header.nb_streamlines as usize,
146 header.nb_vertices as usize,
147 )?
148 }
149 Err(e) => {
150 if header.nb_streamlines == 0 {
151 MmapBacking::OwnedU32(vec![0u32], 4)
152 } else {
153 return Err(e);
154 }
155 }
156 };
157
158 let dps = load_data_dir(&dir.join("dps"))?;
160 let dpv = load_data_dir(&dir.join("dpv"))?;
161 let groups = load_data_dir(&dir.join("groups"))?;
162 let dpg = load_dpg_dir(&dir.join("dpg"))?;
163
164 Ok(TrxFile::from_parts(TrxParts {
165 header,
166 positions_backing,
167 offsets_backing,
168 dps,
169 dpv,
170 groups,
171 dpg,
172 }))
173}
174
175fn convert_offsets_to_u32(
178 mmap: &Mmap,
179 dtype: DType,
180 nb_streamlines: usize,
181 nb_vertices: usize,
182) -> Result<MmapBacking> {
183 match dtype {
184 DType::UInt64 => {
185 let values: &[u64] = cast_slice(mmap.as_ref());
186 if values.len() == nb_streamlines {
188 let mut 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 owned.push(nb_vertices as u32);
199 let bytes: Vec<u8> = crate::mmap_backing::vec_to_bytes(owned);
200 Ok(MmapBacking::Owned(bytes))
201 } else if values.len() == nb_streamlines + 1 {
202 let owned: Vec<u32> = values
203 .iter()
204 .copied()
205 .map(|value| {
206 u32::try_from(value).map_err(|_| {
207 TrxError::Format(format!("offset {value} exceeds uint32 range"))
208 })
209 })
210 .collect::<Result<_>>()?;
211 Ok(MmapBacking::Owned(crate::mmap_backing::vec_to_bytes(owned)))
212 } else {
213 Err(TrxError::Format(format!(
214 "unexpected offset count: {} (expected {} or {})",
215 values.len(),
216 nb_streamlines,
217 nb_streamlines + 1,
218 )))
219 }
220 }
221 DType::UInt32 => {
222 let values: &[u32] = cast_slice(mmap.as_ref());
223 let mut out: Vec<u32> = values.to_vec();
224 if out.len() == nb_streamlines {
225 out.push(nb_vertices as u32);
226 }
227 let bytes: Vec<u8> = crate::mmap_backing::vec_to_bytes(out);
228 Ok(MmapBacking::Owned(bytes))
229 }
230 other => Err(TrxError::DType(format!(
231 "offsets must be uint32 or uint64, got {other}"
232 ))),
233 }
234}
235
236pub fn save_to_directory<P: TrxScalar>(trx: &TrxFile<P>, dir: &Path) -> Result<()> {
239 let offsets_dtype = OffsetsDtype::pick_for(trx.offsets());
240 fs::create_dir_all(dir)?;
241
242 trx.header().write_to(&dir.join("header.json"))?;
244
245 let pos_filename = format!("positions.3.{}", P::DTYPE.name());
247 fs::write(dir.join(&pos_filename), trx.positions_bytes())?;
248
249 let offsets_filename = format!("offsets.{}", offsets_dtype.suffix());
251 let offsets_bytes = offsets_dtype.encode(trx.offsets());
252 fs::write(dir.join(offsets_filename), offsets_bytes)?;
253
254 save_data_dir(trx.dps_arrays(), &dir.join("dps"))?;
256
257 save_data_dir(trx.dpv_arrays(), &dir.join("dpv"))?;
259
260 save_data_dir(trx.group_arrays(), &dir.join("groups"))?;
262
263 save_dpg_dir(trx.dpg_arrays(), &dir.join("dpg"))?;
265
266 Ok(())
267}
268
269pub fn append_dps_to_directory(
271 dir: &Path,
272 dps: &HashMap<String, DataArray>,
273 overwrite: bool,
274) -> Result<()> {
275 let header = Header::from_file(&dir.join("header.json"))?;
276 validate_row_count("DPS", dps, header.nb_streamlines as usize)?;
277 append_arrays_to_directory(&dir.join("dps"), dps, overwrite)
278}
279
280pub fn append_dpv_to_directory(
282 dir: &Path,
283 dpv: &HashMap<String, DataArray>,
284 overwrite: bool,
285) -> Result<()> {
286 let header = Header::from_file(&dir.join("header.json"))?;
287 validate_row_count("DPV", dpv, header.nb_vertices as usize)?;
288 append_arrays_to_directory(&dir.join("dpv"), dpv, overwrite)
289}
290
291pub fn append_groups_to_directory(
293 dir: &Path,
294 groups: &HashMap<String, Vec<u32>>,
295 overwrite: bool,
296) -> Result<()> {
297 let header = Header::from_file(&dir.join("header.json"))?;
298 let groups_dir = dir.join("groups");
299 fs::create_dir_all(&groups_dir)?;
300 for (name, members) in groups {
301 validate_group_members(name, members, header.nb_streamlines as usize)?;
302 let target = groups_dir.join(format!("{name}.uint32"));
303 if !overwrite {
304 if let Some(existing) = find_named_array_file(&groups_dir, name)? {
305 if existing.exists() {
306 continue;
307 }
308 }
309 } else if let Some(existing) = find_named_array_file(&groups_dir, name)? {
310 if existing != target && existing.exists() {
311 fs::remove_file(existing)?;
312 }
313 }
314 fs::write(target, vec_to_bytes(members.clone()))?;
315 }
316 Ok(())
317}
318
319pub fn append_dpg_to_directory(dir: &Path, dpg: &DataPerGroup, overwrite: bool) -> Result<()> {
321 let groups_dir = dir.join("groups");
322 let dpg_root = dir.join("dpg");
323 for (group, entries) in dpg {
324 if find_named_array_file(&groups_dir, group)?.is_none() {
325 return Err(TrxError::Argument(format!(
326 "cannot add DPG entries for missing group '{group}'"
327 )));
328 }
329 let group_dir = dpg_root.join(group);
330 fs::create_dir_all(&group_dir)?;
331 for (name, arr) in entries {
332 let target = group_dir.join(filename_for_array(name, arr));
333 if !overwrite {
334 if let Some(existing) = find_named_array_file(&group_dir, name)? {
335 if existing.exists() {
336 continue;
337 }
338 }
339 } else if let Some(existing) = find_named_array_file(&group_dir, name)? {
340 if existing != target && existing.exists() {
341 fs::remove_file(existing)?;
342 }
343 }
344 fs::write(target, arr.as_bytes())?;
345 }
346 }
347 Ok(())
348}
349
350pub fn delete_dps_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
352 delete_named_arrays(&dir.join("dps"), names)
353}
354
355pub fn delete_dpv_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
357 delete_named_arrays(&dir.join("dpv"), names)
358}
359
360pub fn delete_groups_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
362 let groups_dir = dir.join("groups");
363 for name in names {
364 if let Some(path) = find_named_array_file(&groups_dir, name)? {
365 if path.exists() {
366 fs::remove_file(path)?;
367 }
368 }
369 let dpg_group = dir.join("dpg").join(name);
370 if dpg_group.exists() {
371 fs::remove_dir_all(dpg_group)?;
372 }
373 }
374 Ok(())
375}
376
377pub fn delete_dpg_from_directory(dir: &Path, group: &str, names: Option<&[&str]>) -> Result<()> {
382 let group_dir = dir.join("dpg").join(group);
383 match names {
384 None | Some([]) => {
385 if group_dir.exists() {
386 fs::remove_dir_all(group_dir)?;
387 }
388 }
389 Some(names) => {
390 for name in names {
391 if let Some(path) = find_named_array_file(&group_dir, name)? {
392 if path.exists() {
393 fs::remove_file(path)?;
394 }
395 }
396 }
397 }
398 }
399 Ok(())
400}
401
402fn save_data_dir(arrays: &HashMap<String, DataArray>, dir: &Path) -> Result<()> {
403 if arrays.is_empty() {
404 return Ok(());
405 }
406 fs::create_dir_all(dir)?;
407 for (name, arr) in arrays {
408 let filename = filename_for_array(name, arr);
409 fs::write(dir.join(&filename), arr.as_bytes())?;
410 }
411 Ok(())
412}
413
414fn save_dpg_dir(arrays: &DataPerGroup, dir: &Path) -> Result<()> {
415 if arrays.is_empty() {
416 return Ok(());
417 }
418 fs::create_dir_all(dir)?;
419 for (group, entries) in arrays {
420 save_data_dir(entries, &dir.join(group))?;
421 }
422 Ok(())
423}
424
425fn append_arrays_to_directory(
426 dir: &Path,
427 arrays: &HashMap<String, DataArray>,
428 overwrite: bool,
429) -> Result<()> {
430 fs::create_dir_all(dir)?;
431 for (name, arr) in arrays {
432 let target = dir.join(filename_for_array(name, arr));
433 if !overwrite {
434 if let Some(existing) = find_named_array_file(dir, name)? {
435 if existing.exists() {
436 continue;
437 }
438 }
439 } else if let Some(existing) = find_named_array_file(dir, name)? {
440 if existing != target && existing.exists() {
441 fs::remove_file(existing)?;
442 }
443 }
444 fs::write(target, arr.as_bytes())?;
445 }
446 Ok(())
447}
448
449fn delete_named_arrays(dir: &Path, names: &[&str]) -> Result<()> {
450 for name in names {
451 if let Some(path) = find_named_array_file(dir, name)? {
452 if path.exists() {
453 fs::remove_file(path)?;
454 }
455 }
456 }
457 Ok(())
458}
459
460fn find_named_array_file(dir: &Path, name: &str) -> Result<Option<std::path::PathBuf>> {
461 if !dir.exists() {
462 return Ok(None);
463 }
464 for entry in fs::read_dir(dir)? {
465 let entry = entry?;
466 let path = entry.path();
467 if !path.is_file() {
468 continue;
469 }
470 let file_name = path
471 .file_name()
472 .and_then(|n| n.to_str())
473 .ok_or_else(|| TrxError::Format(format!("invalid filename: {}", path.display())))?;
474 let parsed = TrxFilename::parse(file_name)?;
475 if parsed.name == name {
476 return Ok(Some(path));
477 }
478 }
479 Ok(None)
480}
481
482fn validate_row_count(
483 kind: &str,
484 arrays: &HashMap<String, DataArray>,
485 expected_rows: usize,
486) -> Result<()> {
487 for (name, arr) in arrays {
488 if arr.nrows() != expected_rows {
489 return Err(TrxError::Format(format!(
490 "{kind} '{name}' has {} rows, expected {expected_rows}",
491 arr.nrows()
492 )));
493 }
494 }
495 Ok(())
496}
497
498fn validate_group_members(name: &str, members: &[u32], nb_streamlines: usize) -> Result<()> {
499 for &member in members {
500 if member as usize >= nb_streamlines {
501 return Err(TrxError::Format(format!(
502 "group '{name}' contains streamline index {member}, but NB_STREAMLINES is {nb_streamlines}"
503 )));
504 }
505 }
506 Ok(())
507}
508
509fn filename_for_array(name: &str, arr: &DataArray) -> String {
510 TrxFilename {
511 name: name.to_string(),
512 ncols: arr.ncols(),
513 dtype: arr.dtype(),
514 }
515 .to_filename()
516}