use super::ranges::{convert_signed_index, handle_signed_inclusive_end};
use alloc::vec::Vec;
use core::ops::Range;
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct Slice {
pub start: isize,
pub end: Option<isize>,
pub step: isize,
}
impl Default for Slice {
fn default() -> Self {
Self::full()
}
}
impl Slice {
pub const fn new(start: isize, end: Option<isize>, step: isize) -> Self {
assert!(step != 0, "Step cannot be zero");
Self { start, end, step }
}
pub const fn full() -> Self {
Self::new(0, None, 1)
}
pub fn index(idx: isize) -> Self {
Self {
start: idx,
end: handle_signed_inclusive_end(idx),
step: 1,
}
}
pub fn into_vec(self) -> Vec<isize> {
assert!(
self.end.is_some(),
"Slice must have an end to convert to a vector: {self:?}"
);
self.into_iter().collect()
}
pub fn bound_to(self, size: usize) -> Self {
let mut bounds = size as isize;
if let Some(end) = self.end {
if end > 0 {
bounds = end.min(bounds);
} else {
bounds = end.max(-(bounds + 1));
}
} else if self.is_reversed() {
bounds = -(bounds + 1);
}
Self {
end: Some(bounds),
..self
}
}
pub fn with_step(start: isize, end: Option<isize>, step: isize) -> Self {
assert!(step != 0, "Step cannot be zero");
Self { start, end, step }
}
pub fn from_range_stepped<R: Into<Slice>>(range: R, step: isize) -> Self {
assert!(step != 0, "Step cannot be zero");
let mut slice = range.into();
slice.step = step;
slice
}
pub fn step(&self) -> isize {
self.step
}
pub fn range(&self, size: usize) -> Range<usize> {
self.to_range(size)
}
pub fn to_range(&self, size: usize) -> Range<usize> {
let start = convert_signed_index(self.start, size);
let end = match self.end {
Some(end) => convert_signed_index(end, size),
None => size,
};
start..end
}
pub fn to_range_and_step(&self, size: usize) -> (Range<usize>, isize) {
let range = self.to_range(size);
(range, self.step)
}
pub fn is_reversed(&self) -> bool {
self.step < 0
}
pub fn output_size(&self, dim_size: usize) -> usize {
let range = self.to_range(dim_size);
if range.start >= range.end {
return 0;
}
let len = range.end - range.start;
if self.step.unsigned_abs() == 1 {
len
} else {
len.div_ceil(self.step.unsigned_abs())
}
}
}