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}