1use smallvec::SmallVec;
4
5use std::fmt::Debug;
6use std::ops::{Range, RangeFrom, RangeFull, RangeTo};
7
8#[derive(Clone, Copy, Debug, PartialEq)]
12pub enum SliceItem {
13 Index(isize),
19
20 Range(SliceRange),
22}
23
24impl SliceItem {
25 #[inline]
27 pub fn full_range() -> Self {
28 (..).into()
29 }
30
31 #[inline]
33 pub fn range(start: isize, end: Option<isize>, step: isize) -> SliceItem {
34 SliceItem::Range(SliceRange::new(start, end, step))
35 }
36
37 pub(crate) fn index_range(&self, dim_size: usize) -> IndexRange {
40 let range = match *self {
41 SliceItem::Range(range) => range,
42 SliceItem::Index(idx) => SliceRange::new(idx, Some(idx + 1), 1),
43 };
44 range.index_range(dim_size)
45 }
46}
47
48impl From<i32> for SliceItem {
53 #[inline]
54 fn from(value: i32) -> Self {
55 SliceItem::Index(value as isize)
56 }
57}
58
59impl From<isize> for SliceItem {
60 #[inline]
61 fn from(value: isize) -> Self {
62 SliceItem::Index(value)
63 }
64}
65
66impl From<usize> for SliceItem {
67 #[inline]
68 fn from(value: usize) -> Self {
69 SliceItem::Index(value as isize)
70 }
71}
72
73impl<R> From<R> for SliceItem
74where
75 R: Into<SliceRange>,
76{
77 fn from(value: R) -> Self {
78 SliceItem::Range(value.into())
79 }
80}
81
82pub trait IntoSliceItems {
97 type Array: AsRef<[SliceItem]>;
98
99 fn into_slice_items(self) -> Self::Array;
100}
101
102impl<'a> IntoSliceItems for &'a [SliceItem] {
103 type Array = &'a [SliceItem];
104
105 fn into_slice_items(self) -> &'a [SliceItem] {
106 self
107 }
108}
109
110impl<const N: usize, T: Into<SliceItem>> IntoSliceItems for [T; N] {
111 type Array = [SliceItem; N];
112
113 fn into_slice_items(self) -> [SliceItem; N] {
114 self.map(|x| x.into())
115 }
116}
117
118impl<T: Into<SliceItem>> IntoSliceItems for T {
119 type Array = [SliceItem; 1];
120
121 fn into_slice_items(self) -> [SliceItem; 1] {
122 [self.into()]
123 }
124}
125
126impl<T1: Into<SliceItem>> IntoSliceItems for (T1,) {
127 type Array = [SliceItem; 1];
128
129 fn into_slice_items(self) -> [SliceItem; 1] {
130 [self.0.into()]
131 }
132}
133
134impl<T1: Into<SliceItem>, T2: Into<SliceItem>> IntoSliceItems for (T1, T2) {
135 type Array = [SliceItem; 2];
136
137 fn into_slice_items(self) -> [SliceItem; 2] {
138 [self.0.into(), self.1.into()]
139 }
140}
141
142impl<T1: Into<SliceItem>, T2: Into<SliceItem>, T3: Into<SliceItem>> IntoSliceItems
143 for (T1, T2, T3)
144{
145 type Array = [SliceItem; 3];
146
147 fn into_slice_items(self) -> [SliceItem; 3] {
148 [self.0.into(), self.1.into(), self.2.into()]
149 }
150}
151
152impl<T1: Into<SliceItem>, T2: Into<SliceItem>, T3: Into<SliceItem>, T4: Into<SliceItem>>
153 IntoSliceItems for (T1, T2, T3, T4)
154{
155 type Array = [SliceItem; 4];
156
157 fn into_slice_items(self) -> [SliceItem; 4] {
158 [self.0.into(), self.1.into(), self.2.into(), self.3.into()]
159 }
160}
161
162impl<
163 T1: Into<SliceItem>,
164 T2: Into<SliceItem>,
165 T3: Into<SliceItem>,
166 T4: Into<SliceItem>,
167 T5: Into<SliceItem>,
168> IntoSliceItems for (T1, T2, T3, T4, T5)
169{
170 type Array = [SliceItem; 5];
171
172 fn into_slice_items(self) -> [SliceItem; 5] {
173 [
174 self.0.into(),
175 self.1.into(),
176 self.2.into(),
177 self.3.into(),
178 self.4.into(),
179 ]
180 }
181}
182
183pub type DynSliceItems = SmallVec<[SliceItem; 5]>;
186
187pub fn to_slice_items<T: Clone + Into<SliceItem>>(index: &[T]) -> DynSliceItems {
193 index.iter().map(|x| x.clone().into()).collect()
194}
195
196#[derive(Clone, Copy, Debug, PartialEq)]
208pub struct SliceRange {
209 pub start: isize,
211
212 pub end: Option<isize>,
215
216 step: isize,
219}
220
221impl SliceRange {
222 #[inline]
228 pub fn new(start: isize, end: Option<isize>, step: isize) -> SliceRange {
229 assert!(step != 0, "Slice step cannot be 0");
230 SliceRange { start, end, step }
231 }
232
233 pub fn steps(&self, dim_size: usize) -> usize {
236 let clamped = self.clamp(dim_size);
237
238 let start_idx = Self::offset_from_start(clamped.start, dim_size);
239 let end_idx = clamped
240 .end
241 .map(|index| Self::offset_from_start(index, dim_size))
242 .unwrap_or(if self.step > 0 { dim_size as isize } else { -1 });
243
244 if (clamped.step > 0 && end_idx <= start_idx) || (clamped.step < 0 && end_idx >= start_idx)
245 {
246 return 0;
247 }
248
249 let steps = if clamped.step > 0 {
250 1 + (end_idx - start_idx - 1) / clamped.step
251 } else {
252 1 + (start_idx - end_idx - 1) / -clamped.step
253 };
254
255 steps.max(0) as usize
256 }
257
258 pub fn clamp(&self, dim_size: usize) -> SliceRange {
266 let len = dim_size as isize;
267
268 let min_idx;
269 let max_idx;
270
271 if self.step > 0 {
272 min_idx = -len;
275 max_idx = len;
276 } else {
277 min_idx = -len - 1;
280 max_idx = len - 1;
281 }
282
283 SliceRange::new(
284 self.start.clamp(min_idx, max_idx),
285 self.end.map(|e| e.clamp(min_idx, max_idx)),
286 self.step,
287 )
288 }
289
290 pub fn step(&self) -> isize {
291 self.step
292 }
293
294 pub fn resolve_clamped(&self, dim_size: usize) -> Range<usize> {
300 self.clamp(dim_size).resolve(dim_size).unwrap()
301 }
302
303 #[inline]
311 pub fn resolve(&self, dim_size: usize) -> Option<Range<usize>> {
312 let (start, end) = if self.step > 0 {
313 let start = Self::offset_from_start(self.start, dim_size);
314 let end = self
315 .end
316 .map(|end| Self::offset_from_start(end, dim_size))
317 .unwrap_or(dim_size as isize);
318 (start, end)
319 } else {
320 let start = Self::offset_from_end(self.start, dim_size);
321 let end = self
322 .end
323 .map(|end| Self::offset_from_end(end, dim_size))
324 .unwrap_or(dim_size as isize);
325 (start, end)
326 };
327
328 if start >= 0 && start <= dim_size as isize && end >= 0 && end <= dim_size as isize {
329 let end = end.max(start);
332
333 Some(start as usize..end as usize)
334 } else {
335 None
336 }
337 }
338
339 pub(crate) fn index_range(&self, dim_size: usize) -> IndexRange {
342 let resolved = self.resolve_clamped(dim_size);
345
346 if self.step > 0 {
347 IndexRange::new(resolved.start, resolved.end as isize, self.step)
348 } else {
349 IndexRange::new(
350 dim_size - 1 - resolved.start,
351 dim_size as isize - 1 - resolved.end as isize,
352 self.step,
353 )
354 }
355 }
356
357 #[inline]
359 fn offset_from_start(index: isize, dim_size: usize) -> isize {
360 if index >= 0 {
361 index
362 } else {
363 dim_size as isize + index
364 }
365 }
366
367 #[inline]
369 fn offset_from_end(index: isize, dim_size: usize) -> isize {
370 if index >= 0 {
371 dim_size as isize - 1 - index
372 } else {
373 -index - 1
374 }
375 }
376}
377
378impl<T> From<Range<T>> for SliceRange
379where
380 T: TryInto<isize>,
381 <T as TryInto<isize>>::Error: Debug,
382{
383 fn from(r: Range<T>) -> SliceRange {
384 let start = r.start.try_into().unwrap();
385 let end = r.end.try_into().unwrap();
386 SliceRange::new(start, Some(end), 1)
387 }
388}
389
390impl<T> From<RangeTo<T>> for SliceRange
391where
392 T: TryInto<isize>,
393 <T as TryInto<isize>>::Error: Debug,
394{
395 fn from(r: RangeTo<T>) -> SliceRange {
396 let end = r.end.try_into().unwrap();
397 SliceRange::new(0, Some(end), 1)
398 }
399}
400
401impl<T> From<RangeFrom<T>> for SliceRange
402where
403 T: TryInto<isize>,
404 <T as TryInto<isize>>::Error: Debug,
405{
406 fn from(r: RangeFrom<T>) -> SliceRange {
407 let start = r.start.try_into().unwrap();
408 SliceRange::new(start, None, 1)
409 }
410}
411
412impl From<RangeFull> for SliceRange {
413 #[inline]
414 fn from(_: RangeFull) -> SliceRange {
415 SliceRange::new(0, None, 1)
416 }
417}
418
419#[derive(Copy, Clone, Debug, PartialEq)]
421pub(crate) struct IndexRange {
422 start: usize,
424
425 end: isize,
427 step: isize,
428}
429
430impl IndexRange {
431 fn new(start: usize, end: isize, step: isize) -> Self {
440 assert!(step != 0);
441 assert!(start <= isize::MAX as usize);
442
443 IndexRange {
444 start,
445 end: end.max(-1),
446 step,
447 }
448 }
449
450 #[allow(unused)]
452 pub fn start(&self) -> usize {
453 self.start
454 }
455
456 #[allow(unused)]
459 pub fn end(&self) -> isize {
460 self.end
461 }
462
463 #[allow(unused)]
465 pub fn step(&self) -> isize {
466 self.step
467 }
468
469 pub fn steps(&self) -> usize {
471 let len = if self.step > 0 {
472 (self.end - self.start as isize).max(0).unsigned_abs()
473 } else {
474 (self.end - self.start as isize).min(0).unsigned_abs()
475 };
476 len.div_ceil(self.step.unsigned_abs())
477 }
478}
479
480impl IntoIterator for IndexRange {
481 type Item = usize;
482 type IntoIter = IndexRangeIter;
483
484 #[inline]
485 fn into_iter(self) -> IndexRangeIter {
486 IndexRangeIter {
487 step: self.step,
488 index: self.start as isize,
489 remaining: self.steps(),
490 }
491 }
492}
493
494#[derive(Clone, Debug, PartialEq)]
496pub(crate) struct IndexRangeIter {
497 index: isize,
500
501 remaining: usize,
503
504 step: isize,
505}
506
507impl Iterator for IndexRangeIter {
508 type Item = usize;
509
510 #[inline]
511 fn next(&mut self) -> Option<usize> {
512 if self.remaining == 0 {
513 return None;
514 }
515 let idx = self.index;
516 self.index += self.step;
517 self.remaining -= 1;
518 Some(idx as usize)
519 }
520
521 #[inline]
522 fn size_hint(&self) -> (usize, Option<usize>) {
523 (self.remaining, Some(self.remaining))
524 }
525}
526
527impl ExactSizeIterator for IndexRangeIter {}
528impl std::iter::FusedIterator for IndexRangeIter {}
529
530#[cfg(test)]
531mod tests {
532 use rten_testing::TestCases;
533
534 use super::{IntoSliceItems, SliceItem, SliceRange};
535
536 #[test]
537 fn test_into_slice_items() {
538 let x = (42).into_slice_items();
539 assert_eq!(x, [SliceItem::Index(42)]);
540
541 let x = (2..5).into_slice_items();
542 assert_eq!(x, [SliceItem::Range((2..5).into())]);
543
544 let x = (..5).into_slice_items();
545 assert_eq!(x, [SliceItem::Range((0..5).into())]);
546
547 let x = (3..).into_slice_items();
548 assert_eq!(x, [SliceItem::Range((3..).into())]);
549
550 let x = [1].into_slice_items();
551 assert_eq!(x, [SliceItem::Index(1)]);
552 let x = [1, 2].into_slice_items();
553 assert_eq!(x, [SliceItem::Index(1), SliceItem::Index(2)]);
554
555 let x = (0, 1..2, ..).into_slice_items();
556 assert_eq!(
557 x,
558 [
559 SliceItem::Index(0),
560 SliceItem::Range((1..2).into()),
561 SliceItem::full_range()
562 ]
563 );
564 }
565
566 #[test]
567 fn test_index_range() {
568 #[derive(Debug)]
569 struct Case {
570 range: SliceItem,
571 dim_size: usize,
572 indices: Vec<usize>,
573 }
574
575 let cases = [
576 Case {
578 range: SliceItem::range(0, Some(4), 1),
579 dim_size: 6,
580 indices: (0..4).collect(),
581 },
582 Case {
583 range: SliceItem::range(2, Some(4), 1),
584 dim_size: 6,
585 indices: vec![2, 3],
586 },
587 Case {
588 range: SliceItem::range(2, Some(128), 1),
589 dim_size: 5,
590 indices: vec![2, 3, 4],
591 },
592 Case {
594 range: SliceItem::range(0, Some(5), 2),
595 dim_size: 5,
596 indices: vec![0, 2, 4],
597 },
598 Case {
600 range: SliceItem::range(0, None, 1),
601 dim_size: 6,
602 indices: (0..6).collect(),
603 },
604 Case {
606 range: SliceItem::range(-1, Some(-6), 2),
607 dim_size: 5,
608 indices: vec![],
609 },
610 Case {
612 range: SliceItem::range(-1, Some(-128), -1),
613 dim_size: 5,
614 indices: vec![4, 3, 2, 1, 0],
615 },
616 Case {
618 range: SliceItem::range(-1, None, -1),
619 dim_size: 5,
620 indices: vec![4, 3, 2, 1, 0],
621 },
622 Case {
624 range: SliceItem::range(-1, Some(-6), -2),
625 dim_size: 5,
626 indices: vec![4, 2, 0],
627 },
628 Case {
630 range: SliceItem::range(1, Some(5), -2),
631 dim_size: 5,
632 indices: vec![],
633 },
634 Case {
636 range: SliceItem::range(0, Some(0), 1),
637 dim_size: 4,
638 indices: vec![],
639 },
640 Case {
642 range: SliceItem::range(0, Some(0), -1),
643 dim_size: 4,
644 indices: vec![],
645 },
646 Case {
648 range: SliceItem::Index(2),
649 dim_size: 4,
650 indices: vec![2],
651 },
652 Case {
654 range: SliceItem::Index(2),
655 dim_size: 0,
656 indices: vec![],
657 },
658 ];
659
660 cases.test_each(|case| {
661 let Case {
662 range,
663 dim_size,
664 indices,
665 } = case;
666
667 let mut index_iter = range.index_range(*dim_size).into_iter();
668 let size_hint = index_iter.size_hint();
669 let index_vec: Vec<_> = index_iter.by_ref().collect();
670
671 assert_eq!(size_hint, (index_vec.len(), Some(index_vec.len())));
672 assert_eq!(index_vec, *indices);
673 assert_eq!(index_iter.size_hint(), (0, Some(0)));
674 })
675 }
676
677 #[test]
678 fn test_index_range_steps() {
679 #[derive(Debug)]
680 struct Case {
681 range: SliceRange,
682 dim_size: usize,
683 steps: usize,
684 }
685
686 let cases = [
687 Case {
689 range: SliceRange::new(0, None, 1),
690 dim_size: 4,
691 steps: 4,
692 },
693 Case {
695 range: SliceRange::new(0, None, 5),
696 dim_size: 4,
697 steps: 1,
698 },
699 Case {
701 range: SliceRange::new(-1, None, -1),
702 dim_size: 3,
703 steps: 3,
704 },
705 Case {
707 range: SliceRange::new(1, Some(0), -2),
708 dim_size: 2,
709 steps: 1,
710 },
711 ];
712
713 cases.test_each(|case| {
714 assert_eq!(case.range.index_range(case.dim_size).steps(), case.steps);
715 })
716 }
717
718 #[test]
719 #[should_panic(expected = "Slice step cannot be 0")]
720 fn test_slice_range_zero_step() {
721 SliceRange::new(0, None, 0);
722 }
723
724 #[test]
725 fn test_slice_range_resolve() {
726 assert_eq!(SliceRange::new(0, Some(5), 1).resolve_clamped(10), 0..5);
728 assert_eq!(SliceRange::new(0, None, 1).resolve_clamped(10), 0..10);
729 assert_eq!(SliceRange::new(15, Some(20), 1).resolve_clamped(10), 10..10);
730 assert_eq!(SliceRange::new(15, Some(20), 1).resolve(10), None);
731 assert_eq!(SliceRange::new(4, None, 1).resolve(3), None);
732 assert_eq!(SliceRange::new(0, Some(10), 1).resolve(3), None);
733
734 assert_eq!(SliceRange::new(-5, Some(-1), 1).resolve_clamped(10), 5..9);
736 assert_eq!(SliceRange::new(-20, Some(-1), 1).resolve_clamped(10), 0..9);
737 assert_eq!(SliceRange::new(-20, Some(-1), 1).resolve(10), None);
738 assert_eq!(SliceRange::new(-5, None, 1).resolve_clamped(10), 5..10);
739
740 assert_eq!(SliceRange::new(5, Some(0), -1).resolve_clamped(10), 4..9);
745 assert_eq!(SliceRange::new(5, None, -1).resolve_clamped(10), 4..10);
746 assert_eq!(SliceRange::new(9, None, -1).resolve_clamped(10), 0..10);
747
748 assert_eq!(SliceRange::new(-1, Some(-4), -1).resolve_clamped(3), 0..3);
750 assert_eq!(SliceRange::new(-1, None, -1).resolve_clamped(2), 0..2);
751 }
752}