Skip to main content

trx_rs/io/
filename.rs

1use crate::dtype::DType;
2use crate::error::{Result, TrxError};
3
4/// Parsed components of a TRX data filename like `positions.3.float32`.
5#[derive(Debug, Clone, PartialEq, Eq)]
6pub struct TrxFilename {
7    pub name: String,
8    pub ncols: usize,
9    pub dtype: DType,
10}
11
12impl TrxFilename {
13    /// Parse a filename stem (no directory, no leading path).
14    ///
15    /// Accepts two formats:
16    /// - `{name}.{ncols}.{dtype}` (e.g. `positions.3.float32`)
17    /// - `{name}.{dtype}` (e.g. `offsets.uint32`) — ncols defaults to 1
18    pub fn parse(stem: &str) -> Result<Self> {
19        // Try 3-part format first: split from the right.
20        let parts3: Vec<&str> = stem.rsplitn(3, '.').collect();
21
22        if parts3.len() == 3 {
23            // Could be {name}.{ncols}.{dtype} or {name}.{something}.{dtype}
24            if let Ok(ncols) = parts3[1].parse::<usize>() {
25                if let Ok(dtype) = DType::parse(parts3[0]) {
26                    return Ok(TrxFilename {
27                        name: parts3[2].to_string(),
28                        ncols,
29                        dtype,
30                    });
31                }
32            }
33        }
34
35        // Try 2-part format: {name}.{dtype} (ncols = 1)
36        let parts2: Vec<&str> = stem.rsplitn(2, '.').collect();
37        if parts2.len() == 2 {
38            if let Ok(dtype) = DType::parse(parts2[0]) {
39                return Ok(TrxFilename {
40                    name: parts2[1].to_string(),
41                    ncols: 1,
42                    dtype,
43                });
44            }
45        }
46
47        Err(TrxError::Format(format!(
48            "cannot parse TRX filename '{stem}'"
49        )))
50    }
51
52    /// Format back to `{name}.{dtype}` for 1D arrays or `{name}.{ncols}.{dtype}` otherwise.
53    pub fn to_filename(&self) -> String {
54        if self.ncols == 1 {
55            format!("{}.{}", self.name, self.dtype.name())
56        } else {
57            format!("{}.{}.{}", self.name, self.ncols, self.dtype.name())
58        }
59    }
60}
61
62#[cfg(test)]
63mod tests {
64    use super::*;
65
66    #[test]
67    fn parse_positions() {
68        let f = TrxFilename::parse("positions.3.float32").unwrap();
69        assert_eq!(f.name, "positions");
70        assert_eq!(f.ncols, 3);
71        assert_eq!(f.dtype, DType::Float32);
72    }
73
74    #[test]
75    fn parse_offsets() {
76        let f = TrxFilename::parse("offsets.1.uint32").unwrap();
77        assert_eq!(f.name, "offsets");
78        assert_eq!(f.ncols, 1);
79        assert_eq!(f.dtype, DType::UInt32);
80    }
81
82    #[test]
83    fn round_trip() {
84        let f = TrxFilename {
85            name: "fa".into(),
86            ncols: 1,
87            dtype: DType::Float32,
88        };
89        assert_eq!(f.to_filename(), "fa.float32");
90        assert_eq!(TrxFilename::parse(&f.to_filename()).unwrap(), f);
91    }
92
93    #[test]
94    fn parse_two_part() {
95        // 1D data elements don't include ncols — ncols defaults to 1
96        let f = TrxFilename::parse("offsets.uint32").unwrap();
97        assert_eq!(f.name, "offsets");
98        assert_eq!(f.ncols, 1);
99        assert_eq!(f.dtype, DType::UInt32);
100    }
101
102    #[test]
103    fn parse_invalid() {
104        assert!(TrxFilename::parse("noext").is_err());
105        assert!(TrxFilename::parse("foo.bar").is_err());
106    }
107
108    #[test]
109    fn name_with_dots() {
110        let f = TrxFilename::parse("my.metric.1.float64").unwrap();
111        assert_eq!(f.name, "my.metric");
112        assert_eq!(f.ncols, 1);
113        assert_eq!(f.dtype, DType::Float64);
114    }
115}