1use itertools::izip;
2use serde::{Deserialize, Serialize};
3
4use crate::circuit::errors::SliceError;
5
6#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
9pub struct Slice(SliceEnum);
10
11#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, Hash, Serialize, Deserialize)]
12#[repr(C)]
13enum SliceEnum {
14 Single(u32),
16 Range { start: u32, size: u32, step: i64 },
21 Range2d {
26 start: u32,
27 size1: u32,
28 step1: i64,
29 size2: u32,
30 step2: i64,
31 },
32 RangeVec(Vec<SliceEnum>),
34}
35
36impl Slice {
37 pub fn empty() -> Self {
38 Self(SliceEnum::RangeVec(vec![]))
39 }
40
41 pub fn single(index: u32) -> Self {
42 Self(SliceEnum::Single(index))
43 }
44
45 pub fn range(start: u32, size: u32, step: i64) -> Result<Self, SliceError> {
46 validate_range_bounds(start, size, step)?;
47 Ok(Self(SliceEnum::Range { start, size, step }))
48 }
49
50 pub fn shift_start(&mut self, delta: u32) {
51 self.0.shift_start(delta);
52 }
53
54 pub fn range2d(
55 start: u32,
56 size1: u32,
57 size2: u32,
58 step1: i64,
59 step2: i64,
60 ) -> Result<Self, SliceError> {
61 validate_range_2d_bounds(start, size1, step1, size2, step2)?;
62 Ok(Self(SliceEnum::Range2d {
63 start,
64 size1,
65 step1,
66 size2,
67 step2,
68 }))
69 }
70
71 pub fn append(&mut self, other: Self) {
72 match (&mut self.0, other.0) {
73 (SliceEnum::RangeVec(v), SliceEnum::RangeVec(v1)) => v.extend(v1),
74 (SliceEnum::RangeVec(v), slice) => v.push(slice),
75 (slice, SliceEnum::RangeVec(mut v1)) => {
76 v1.insert(0, slice.clone());
77 *slice = SliceEnum::RangeVec(v1);
78 }
79 (slice, slice1) => *slice = SliceEnum::RangeVec(vec![slice.clone(), slice1]),
80 }
81 }
82
83 pub fn get_indices(&self) -> Vec<u32> {
84 self.0.get_indices()
85 }
86
87 pub fn is_empty(&self) -> bool {
88 self.len() == 0
89 }
90
91 pub fn len(&self) -> u32 {
92 self.0.len()
93 }
94
95 pub fn from_indices(indices: Vec<u32>) -> Self {
96 Self(SliceEnum::from_indices(indices))
97 }
98
99 pub fn optimize(self) -> Self {
100 Self::from_indices(self.get_indices())
101 }
102}
103
104fn validate_bounds(min_index: i128, max_index: i128) -> Result<(), SliceError> {
105 if min_index < 0 {
106 return Err(SliceError::NegativeIndex(min_index));
107 }
108 if max_index > i128::from(u32::MAX) {
109 return Err(SliceError::IndexOutOfBounds {
110 found: max_index,
111 max: u32::MAX,
112 });
113 }
114 Ok(())
115}
116
117#[inline]
118fn range_index(start: u32, step: i64, i: i64) -> i128 {
119 i128::from(start) + i128::from(step) * i128::from(i)
120}
121
122#[inline]
123fn range_2d_index(start: u32, step1: i64, i: i64, step2: i64, j: i64) -> i128 {
124 i128::from(start) + i128::from(step1) * i128::from(i) + i128::from(step2) * i128::from(j)
125}
126
127fn validate_range_bounds(start: u32, size: u32, step: i64) -> Result<(), SliceError> {
128 if size == 0 {
129 return Ok(());
130 }
131 let last = i64::from(size - 1);
132 let first = i128::from(start);
133 let end = range_index(start, step, last);
134 validate_bounds(first.min(end), first.max(end))
135}
136
137fn validate_range_2d_bounds(
138 start: u32,
139 size1: u32,
140 step1: i64,
141 size2: u32,
142 step2: i64,
143) -> Result<(), SliceError> {
144 if size1 == 0 || size2 == 0 {
145 return Ok(());
146 }
147 let i_last = i64::from(size1 - 1);
148 let j_last = i64::from(size2 - 1);
149
150 let corners = [
151 range_2d_index(start, step1, 0, step2, 0),
152 range_2d_index(start, step1, i_last, step2, 0),
153 range_2d_index(start, step1, 0, step2, j_last),
154 range_2d_index(start, step1, i_last, step2, j_last),
155 ];
156
157 let min_index = corners.into_iter().min().unwrap_or(0);
158 let max_index = corners.into_iter().max().unwrap_or(0);
159 validate_bounds(min_index, max_index)
160}
161
162#[inline]
163fn to_u32_index(index: i128) -> u32 {
164 u32::try_from(index).unwrap_or_else(|_| panic!("slice index out of bounds: {index}"))
165}
166
167fn generate_range_indices(start: u32, size: u32, step: i64) -> impl Iterator<Item = u32> {
168 (0..i64::from(size)).map(move |i| to_u32_index(range_index(start, step, i)))
169}
170
171fn generate_range_2d_indices(
172 start: u32,
173 size1: u32,
174 step1: i64,
175 size2: u32,
176 step2: i64,
177) -> impl Iterator<Item = u32> {
178 (0..i64::from(size1)).flat_map(move |i| {
179 (0..i64::from(size2)).map(move |j| to_u32_index(range_2d_index(start, step1, i, step2, j)))
180 })
181}
182
183impl SliceEnum {
184 fn get_indices(&self) -> Vec<u32> {
185 match self {
186 SliceEnum::Single(idx) => vec![*idx],
187 SliceEnum::Range { start, size, step } => {
188 generate_range_indices(*start, *size, *step).collect()
189 }
190 SliceEnum::Range2d {
191 start,
192 size1,
193 size2,
194 step1,
195 step2,
196 } => generate_range_2d_indices(*start, *size1, *step1, *size2, *step2).collect(),
197 SliceEnum::RangeVec(v) => v.iter().flat_map(|r| r.get_indices()).collect(),
198 }
199 }
200
201 pub fn len(&self) -> u32 {
202 match self {
203 SliceEnum::Single(_) => 1,
204 SliceEnum::Range { size, .. } => *size,
205 SliceEnum::Range2d { size1, size2, .. } => size1
206 .checked_mul(*size2)
207 .expect("slice length overflow for range2d"),
208 SliceEnum::RangeVec(v) => v.iter().fold(0u32, |acc, r| {
209 acc.checked_add(r.len())
210 .expect("slice length overflow for range vector")
211 }),
212 }
213 }
214
215 fn match_largest_slice(start: u32, deltas: &[i64]) -> Self {
219 if deltas.is_empty() {
220 return Self::Single(start);
221 }
222
223 let step_j = deltas[0];
226 let n_j = deltas.iter().skip(1).take_while(|&&d| d == step_j).count() + 2;
227
228 let mut res_slice = Self::Range {
229 start,
230 size: n_j as u32,
231 step: step_j,
232 };
233
234 if n_j < deltas.len() + 1 {
235 let exp_chunk = &deltas[0..n_j];
239 let chunks = deltas.chunks(n_j).skip(1);
240 let mut n_i = chunks
241 .take_while(|chunk| {
242 izip!(exp_chunk, *chunk).take_while(|(e, d)| e == d).count() == n_j
243 })
244 .count()
245 + 1;
246 if let Some(chunk) = deltas.chunks(n_j).nth(n_i) {
247 if izip!(exp_chunk, chunk).take_while(|(e, d)| e == d).count() == n_j - 1 {
248 n_i += 1;
249 }
250 }
251
252 if n_i > 1 {
253 let step_i = exp_chunk.iter().sum::<i64>();
254 res_slice = Self::Range2d {
255 start,
256 size1: n_i as u32,
257 size2: n_j as u32,
258 step1: step_i,
259 step2: step_j,
260 };
261 }
262 }
263
264 res_slice
265 }
266
267 fn reduce(&mut self, max_size: u32) {
269 assert!(max_size > 0);
270 match self {
271 SliceEnum::Single(_) => {}
272 SliceEnum::Range { start, size, .. } => {
273 if max_size < *size {
274 if max_size == 1 {
275 *self = SliceEnum::Single(*start);
276 } else {
277 *size = max_size;
278 }
279 }
280 }
281 SliceEnum::Range2d {
282 start,
283 size1,
284 size2,
285 step2,
286 ..
287 } => {
288 if max_size < *size1 * *size2 {
289 if max_size == 1 {
290 *self = SliceEnum::Single(*start);
291 } else if max_size <= *size2 {
292 *self = SliceEnum::Range {
293 start: *start,
294 size: max_size,
295 step: *step2,
296 }
297 } else if max_size / *size2 == 1 {
298 *self = SliceEnum::Range {
299 start: *start,
300 size: *size2,
301 step: *step2,
302 }
303 } else {
304 *size1 = max_size / *size2;
305 }
306 }
307 }
308 SliceEnum::RangeVec(_) => {}
309 }
310 }
311
312 fn match_slices(mut max_len_slices: Vec<Self>) -> Vec<Self> {
313 let mut res = vec![]; let mut ranges_to_visit = vec![(0, max_len_slices.len())]; while let Some((start, end)) = ranges_to_visit.pop() {
316 let (slice_pos, slice) = max_len_slices[start..end]
319 .iter()
320 .enumerate()
321 .max_by_key(|(pos, slice)| (slice.len(), end - pos)) .unwrap();
323 let slice_start = start + slice_pos; let slice_end = slice_start + slice.len() as usize;
325
326 res.push((slice_start, slice.clone()));
328
329 if start < slice_start {
331 max_len_slices[start..slice_start]
334 .iter_mut()
335 .enumerate()
336 .for_each(|(pos, slice)| slice.reduce((slice_pos - pos) as u32));
337
338 ranges_to_visit.push((start, slice_start));
339 }
340 if slice_end < end {
341 ranges_to_visit.push((slice_end, end));
342 }
343 }
344
345 res.sort_by_key(|(start, _)| *start);
346 res.into_iter().map(|(_, slice)| slice).collect()
347 }
348
349 pub fn from_indices(indices: Vec<u32>) -> Self {
356 if indices.is_empty() {
357 return Self::RangeVec(vec![]);
358 }
359
360 let deltas = indices
361 .windows(2)
362 .map(|w| w[1] as i64 - w[0] as i64)
363 .collect::<Vec<_>>();
364 let max_slice_vec: Vec<_> = (0..indices.len())
365 .map(|i| Self::match_largest_slice(indices[i], &deltas[i..]))
366 .collect();
367
368 let optimized_slices = SliceEnum::match_slices(max_slice_vec);
369 if optimized_slices.len() == 1 {
370 optimized_slices[0].clone()
371 } else {
372 SliceEnum::RangeVec(optimized_slices)
373 }
374 }
375
376 pub fn shift_start(&mut self, delta: u32) {
377 match self {
378 SliceEnum::Single(idx) => {
379 *idx = idx
380 .checked_add(delta)
381 .expect("slice start overflow for single index");
382 }
383 SliceEnum::Range { start, .. } => {
384 *start = start
385 .checked_add(delta)
386 .expect("slice start overflow for range");
387 }
388 SliceEnum::Range2d { start, .. } => {
389 *start = start
390 .checked_add(delta)
391 .expect("slice start overflow for range2d");
392 }
393 SliceEnum::RangeVec(v) => v.iter_mut().for_each(|slice| slice.shift_start(delta)),
394 }
395 }
396}
397
398#[cfg(test)]
399mod tests {
400 use super::SliceEnum;
401 use crate::circuit::{errors::SliceError, Slice};
402
403 #[test]
404 fn test_slice_range() {
405 let range = SliceEnum::Range2d {
406 start: 0,
407 size1: 2,
408 size2: 3,
409 step1: 6,
410 step2: 1,
411 };
412 let expected = vec![0, 1, 2, 6, 7, 8];
413 assert_eq!(range.get_indices(), expected);
414
415 let range = SliceEnum::Range2d {
416 start: 0,
417 size1: 4,
418 size2: 2,
419 step1: 3,
420 step2: 1,
421 };
422 let expected = vec![0, 1, 3, 4, 6, 7, 9, 10];
423 assert_eq!(range.get_indices(), expected);
424
425 let range = SliceEnum::Range2d {
426 start: 0,
427 size1: 4,
428 size2: 2,
429 step1: 3,
430 step2: 2,
431 };
432 let expected = vec![0, 2, 3, 5, 6, 8, 9, 11];
433 assert_eq!(range.get_indices(), expected);
434
435 let range = SliceEnum::Range2d {
436 start: 2,
437 size1: 1,
438 size2: 4,
439 step1: 1,
440 step2: 3,
441 };
442 let expected = vec![2, 5, 8, 11];
443 assert_eq!(range.get_indices(), expected);
444 }
445
446 #[test]
447 fn test_slice_match_largest_slice() {
448 fn match_largest_slice(indices: &[u32]) -> SliceEnum {
449 SliceEnum::match_largest_slice(
450 indices[0],
451 &indices
452 .windows(2)
453 .map(|w| w[1] as i64 - w[0] as i64)
454 .collect::<Vec<_>>(),
455 )
456 }
457
458 let indices = vec![0];
461 let slice = match_largest_slice(&indices);
462 assert_eq!(slice.get_indices(), indices);
463
464 let indices = vec![3];
465 let slice = match_largest_slice(&indices);
466 assert_eq!(slice.get_indices(), indices);
467
468 let indices = vec![0, 1, 2, 3, 4];
470 let slice = match_largest_slice(&indices);
471 assert_eq!(slice.get_indices(), indices);
472
473 let indices = vec![5, 7, 9, 11, 13];
474 let slice = match_largest_slice(&indices);
475 assert_eq!(slice.get_indices(), indices);
476
477 let indices = vec![5, 6];
478 let slice = match_largest_slice(&indices);
479 assert_eq!(slice.get_indices(), indices);
480
481 let indices = vec![5, 2];
482 let slice = match_largest_slice(&indices);
483 assert_eq!(slice.get_indices(), indices[..2].to_vec());
484
485 let indices = vec![0, 1, 2, 5, 6, 7, 10, 11, 12, 15, 16, 17]; let slice = match_largest_slice(&indices);
488 assert_eq!(slice.get_indices(), indices);
489
490 let indices = vec![2, 3, 4, 7, 8, 9]; let slice = match_largest_slice(&indices);
492 assert_eq!(slice.get_indices(), indices);
493
494 let indices = vec![0, 2, 8, 10]; let slice = match_largest_slice(&indices);
496 assert_eq!(slice.get_indices(), indices);
497
498 let indices = vec![10, 12, 5, 7, 0, 2]; let slice = match_largest_slice(&indices);
500 assert_eq!(slice.get_indices(), indices.to_vec());
501
502 let indices = vec![0, 2, 4, 4, 5];
505 let slice = match_largest_slice(&indices);
506 assert_eq!(slice.get_indices(), indices[..3].to_vec());
507
508 let indices = vec![0, 1, 3, 4, 5];
510 let slice = match_largest_slice(&indices);
511 assert_eq!(slice.get_indices(), indices[..4].to_vec());
512
513 let indices = vec![10, 12, 5, 7, 0, 2, 1];
514 let slice = match_largest_slice(&indices);
515 assert_eq!(slice.get_indices(), indices[..6].to_vec());
516
517 let indices = vec![1, 1, 0, 0, 1, 1, 0, 0];
519 let slice = match_largest_slice(&indices);
520 assert_eq!(slice.get_indices(), indices[..4].to_vec());
521 }
522
523 #[test]
524 fn test_slice_optimize() {
525 let indices = vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11];
526 let slice = Slice::from_indices(indices.clone());
527 assert_eq!(slice.get_indices(), indices);
528 assert_eq!(
529 slice.0,
530 SliceEnum::Range {
531 start: 0,
532 size: 12,
533 step: 1,
534 }
535 );
536
537 let indices = vec![19, 3, 4, 5, 6, 7, 8, 9, 10, 11];
538 let slice = Slice::from_indices(indices.clone());
539 assert_eq!(slice.get_indices(), indices);
540 assert_eq!(
541 slice.0,
542 SliceEnum::RangeVec(vec![
543 SliceEnum::Single(19),
544 SliceEnum::Range {
545 start: 3,
546 size: 9,
547 step: 1
548 }
549 ])
550 );
551
552 let indices = vec![0, 1, 2, 19, 3, 4, 5, 6, 7, 8, 9, 10, 11];
553 let slice = Slice::from_indices(indices.clone());
554 assert_eq!(slice.get_indices(), indices);
555 assert_eq!(
556 slice.0,
557 SliceEnum::RangeVec(vec![
558 SliceEnum::Range {
559 start: 0,
560 size: 3,
561 step: 1
562 },
563 SliceEnum::Single(19),
564 SliceEnum::Range {
565 start: 3,
566 size: 9,
567 step: 1
568 }
569 ])
570 );
571
572 let indices = vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 19];
573 let slice = Slice::from_indices(indices.clone());
574 assert_eq!(slice.get_indices(), indices);
575 assert_eq!(
576 slice.0,
577 SliceEnum::RangeVec(vec![
578 SliceEnum::Range {
579 start: 0,
580 size: 10,
581 step: 1
582 },
583 SliceEnum::Single(19),
584 ])
585 );
586
587 let indices = vec![0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 19, 10, 11];
588 let slice = Slice::from_indices(indices.clone());
589 assert_eq!(slice.get_indices(), indices);
590 assert_eq!(
591 slice.0,
592 SliceEnum::RangeVec(vec![
593 SliceEnum::Range {
594 start: 0,
595 size: 10,
596 step: 1
597 },
598 SliceEnum::Range {
599 start: 19,
600 size: 2,
601 step: -9
602 },
603 SliceEnum::Single(11),
604 ])
605 );
606
607 let mut indices = Vec::new();
609 for _i in 0..1000 {
610 indices.extend(vec![0, 1, 1, 0]);
611 }
612 let slice = Slice::from_indices(indices.clone());
613 assert_eq!(slice.get_indices(), indices);
614 }
615
616 #[test]
617 fn test_slice_checked_range_bounds() {
618 assert_eq!(Slice::range(0, 2, -1), Err(SliceError::NegativeIndex(-1)));
619 assert_eq!(
620 Slice::range(u32::MAX, 2, 1),
621 Err(SliceError::IndexOutOfBounds {
622 found: i128::from(u32::MAX) + 1,
623 max: u32::MAX
624 })
625 );
626
627 let slice = Slice::range(u32::MAX - 1, 2, 1).unwrap();
628 assert_eq!(slice.get_indices(), vec![u32::MAX - 1, u32::MAX]);
629 }
630
631 #[test]
632 fn test_slice_checked_range2d_bounds() {
633 assert_eq!(
634 Slice::range2d(0, 2, 2, -1, 0),
635 Err(SliceError::NegativeIndex(-1))
636 );
637 assert_eq!(
638 Slice::range2d(u32::MAX, 2, 1, 1, 0),
639 Err(SliceError::IndexOutOfBounds {
640 found: i128::from(u32::MAX) + 1,
641 max: u32::MAX
642 })
643 );
644
645 let slice = Slice::range2d(u32::MAX - 1, 1, 2, 1, 1).unwrap();
646 assert_eq!(slice.get_indices(), vec![u32::MAX - 1, u32::MAX]);
647 }
648}