Skip to main content

trx_rs/io/
directory.rs

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
15// Re-use `OffsetsDtype` so the directory writer matches the zip writer.
16pub(crate) use super::zip::OffsetsDtype;
17
18/// Memory-map a file as read-only.
19fn mmap_file(path: &Path) -> Result<Mmap> {
20    let file = fs::File::open(path)?;
21    // SAFETY: We trust the file won't be modified externally while mapped.
22    let mmap = unsafe { Mmap::map(&file)? };
23    Ok(mmap)
24}
25
26/// Load data arrays from a subdirectory (e.g. `dps/`, `dpv/`, `groups/`).
27fn 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
78/// Find a file matching a given name prefix in a directory, regardless of
79/// the ncols/dtype suffix.
80fn 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
92/// Load a `TrxFile<P>` from an uncompressed directory.
93pub 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    // Header
99    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    // DPS, DPV, groups
159    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
175/// Convert offset bytes to u32, handling uint64→u32 narrowing and
176/// ensuring the sentinel value (nb_vertices) is present.
177fn 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            // Check if sentinel is present
187            if values.len() == nb_streamlines {
188                // Missing sentinel — append nb_vertices
189                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
236/// Save a `TrxFile<P>` to an uncompressed directory. The `offsets.*` array
237/// width is auto-picked: `uint32` when every offset fits, otherwise `uint64`.
238pub 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    // Header
243    trx.header().write_to(&dir.join("header.json"))?;
244
245    // Positions
246    let pos_filename = format!("positions.3.{}", P::DTYPE.name());
247    fs::write(dir.join(&pos_filename), trx.positions_bytes())?;
248
249    // Offsets — written at `offsets_dtype`'s width.
250    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    // DPS
255    save_data_dir(trx.dps_arrays(), &dir.join("dps"))?;
256
257    // DPV
258    save_data_dir(trx.dpv_arrays(), &dir.join("dpv"))?;
259
260    // Groups
261    save_data_dir(trx.group_arrays(), &dir.join("groups"))?;
262
263    // DPG
264    save_dpg_dir(trx.dpg_arrays(), &dir.join("dpg"))?;
265
266    Ok(())
267}
268
269/// Append DPS arrays to a TRX directory, optionally overwriting existing entries.
270pub 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
280/// Append DPV arrays to a TRX directory, optionally overwriting existing entries.
281pub 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
291/// Append group membership arrays to a TRX directory, optionally overwriting existing entries.
292pub 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
319/// Append DPG (data-per-group) entries to a TRX directory, optionally overwriting existing entries.
320pub 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
350/// Delete named DPS arrays from a TRX directory.
351pub fn delete_dps_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
352    delete_named_arrays(&dir.join("dps"), names)
353}
354
355/// Delete named DPV arrays from a TRX directory.
356pub fn delete_dpv_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
357    delete_named_arrays(&dir.join("dpv"), names)
358}
359
360/// Delete named groups (and their DPG entries) from a TRX directory.
361pub 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
377/// Delete DPG entries for a specific group from a TRX directory.
378///
379/// When `names` is `None` or empty, the entire DPG subdirectory for the group is removed.
380/// When `names` lists specific fields, only those entries are deleted.
381pub 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}