Skip to main content

tenferro_ops/
axis.rs

1use thiserror::Error;
2
3/// Error returned when normalizing user-facing axis arguments.
4#[derive(Debug, Clone, PartialEq, Eq, Error)]
5pub enum AxisError {
6    /// Axis is outside `[-rank, rank)`.
7    #[error("axis {axis} is out of bounds for rank {rank}")]
8    OutOfBounds { axis: isize, rank: usize },
9    /// Axis appears more than once after negative-axis normalization.
10    #[error("duplicate axis {axis}")]
11    Duplicate { axis: usize },
12}
13
14/// Normalize a possibly-negative axis against `rank`.
15///
16/// # Examples
17///
18/// ```
19/// use tenferro_ops::axis::normalize_axis;
20///
21/// assert_eq!(normalize_axis(-1, 3).unwrap(), 2);
22/// assert!(normalize_axis(3, 3).is_err());
23/// ```
24///
25/// # Errors
26///
27/// Returns [`AxisError::OutOfBounds`] when `axis` is outside the normalized
28/// range for `rank`.
29pub fn normalize_axis(axis: isize, rank: usize) -> Result<usize, AxisError> {
30    let normalized = if axis >= 0 {
31        axis as usize
32    } else {
33        rank.checked_sub(axis.unsigned_abs())
34            .ok_or(AxisError::OutOfBounds { axis, rank })?
35    };
36    if normalized >= rank {
37        return Err(AxisError::OutOfBounds { axis, rank });
38    }
39    Ok(normalized)
40}
41
42/// Normalize a list of possibly-negative axes and reject duplicates.
43///
44/// # Examples
45///
46/// ```
47/// use tenferro_ops::axis::normalize_axes;
48///
49/// assert_eq!(normalize_axes(&[0, -1], 3).unwrap(), vec![0, 2]);
50/// assert!(normalize_axes(&[1, -2], 3).is_err());
51/// ```
52///
53/// # Errors
54///
55/// Returns [`AxisError::OutOfBounds`] for an invalid axis or
56/// [`AxisError::Duplicate`] when two axes normalize to the same position.
57pub fn normalize_axes(axes: &[isize], rank: usize) -> Result<Vec<usize>, AxisError> {
58    let mut out = Vec::with_capacity(axes.len());
59    let mut seen = vec![false; rank];
60    for &axis in axes {
61        let normalized = normalize_axis(axis, rank)?;
62        if seen[normalized] {
63            return Err(AxisError::Duplicate { axis: normalized });
64        }
65        seen[normalized] = true;
66        out.push(normalized);
67    }
68    Ok(out)
69}