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>(
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    // Header
102    let header = Header::from_file(&dir.join("header.json"))?;
103
104    // Positions
105    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    // Offsets
129    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    // DPS, DPV, groups
145    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
162/// Convert offset bytes to u32, handling uint64→u32 narrowing and
163/// ensuring the sentinel value (nb_vertices) is present.
164fn 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            // Check if sentinel is present
174            if values.len() == nb_streamlines {
175                // Missing sentinel — append nb_vertices
176                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
223/// Save a `TrxFile<P>` to an uncompressed directory. The `offsets.*` array
224/// width is auto-picked: `uint32` when every offset fits, otherwise `uint64`.
225pub 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    // Header
230    trx.header().write_to(&dir.join("header.json"))?;
231
232    // Positions
233    let pos_filename = format!("positions.3.{}", P::DTYPE.name());
234    fs::write(dir.join(&pos_filename), trx.positions_bytes())?;
235
236    // Offsets — written at `offsets_dtype`'s width.
237    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    // DPS
242    save_data_dir(trx.dps_arrays(), &dir.join("dps"))?;
243
244    // DPV
245    save_data_dir(trx.dpv_arrays(), &dir.join("dpv"))?;
246
247    // Groups
248    save_data_dir(trx.group_arrays(), &dir.join("groups"))?;
249
250    // DPG
251    save_dpg_dir(trx.dpg_arrays(), &dir.join("dpg"))?;
252
253    Ok(())
254}
255
256/// Append DPS arrays to a TRX directory, optionally overwriting existing entries.
257pub 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
267/// Append DPV arrays to a TRX directory, optionally overwriting existing entries.
268pub 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
278/// Append group membership arrays to a TRX directory, optionally overwriting existing entries.
279pub 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
306/// Append DPG (data-per-group) entries to a TRX directory, optionally overwriting existing entries.
307pub 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
337/// Delete named DPS arrays from a TRX directory.
338pub fn delete_dps_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
339    delete_named_arrays(&dir.join("dps"), names)
340}
341
342/// Delete named DPV arrays from a TRX directory.
343pub fn delete_dpv_from_directory(dir: &Path, names: &[&str]) -> Result<()> {
344    delete_named_arrays(&dir.join("dpv"), names)
345}
346
347/// Delete named groups (and their DPG entries) from a TRX directory.
348pub 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
364/// Delete DPG entries for a specific group from a TRX directory.
365///
366/// When `names` is `None` or empty, the entire DPG subdirectory for the group is removed.
367/// When `names` lists specific fields, only those entries are deleted.
368pub 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}