Skip to main content

ferritin_core/data/
segmentation.rs

1//! CSR-style segmentation primitive for grouping elements into segments.
2//!
3//! `Segmentation` stores an offset array where `offsets[i]..offsets[i+1]`
4//! is the element range for segment `i`.  This is the same layout as a
5//! Compressed Sparse Row (CSR) row-pointer array.
6
7use std::ops::Range;
8
9/// CSR-style segmentation: groups elements into segments via offset array.
10///
11/// `offsets.len() == n_segments + 1`; segment `i` spans
12/// `offsets[i]..offsets[i+1]`.
13#[derive(Clone, Debug, PartialEq)]
14pub struct Segmentation {
15    offsets: Vec<u32>, // len = n_segments + 1
16}
17
18impl Segmentation {
19    /// Build from pre-computed offsets (e.g., `[0, 5, 12, 20]` for 3 segments).
20    ///
21    /// # Panics
22    /// Panics if `offsets` is empty (needs at least the sentinel `[0]`).
23    pub fn from_offsets(offsets: Vec<u32>) -> Self {
24        assert!(!offsets.is_empty(), "offsets must have at least one element");
25        Self { offsets }
26    }
27
28    /// Build by detecting change-points in a key sequence.
29    ///
30    /// Each run of equal consecutive keys becomes one segment.
31    ///
32    /// # Example
33    /// ```
34    /// # use ferritin_core::data::Segmentation;
35    /// let seg = Segmentation::from_change_points(["A","A","A","B","B","C"].iter().copied());
36    /// assert_eq!(seg.count(), 3);
37    /// ```
38    pub fn from_change_points<K: PartialEq>(keys: impl Iterator<Item = K>) -> Self {
39        let mut offsets: Vec<u32> = vec![0];
40        let mut prev: Option<K> = None;
41        let mut idx: u32 = 0;
42
43        for key in keys {
44            match &prev {
45                None => {}
46                Some(p) if *p == key => {}
47                Some(_) => {
48                    offsets.push(idx);
49                }
50            }
51            prev = Some(key);
52            idx += 1;
53        }
54        offsets.push(idx); // final sentinel
55        Self { offsets }
56    }
57
58    /// Number of segments.
59    pub fn count(&self) -> usize {
60        self.offsets.len().saturating_sub(1)
61    }
62
63    /// Element range for segment `seg`.
64    ///
65    /// # Panics
66    /// Panics if `seg >= self.count()`.
67    pub fn segment(&self, seg: usize) -> Range<usize> {
68        let start = self.offsets[seg] as usize;
69        let end = self.offsets[seg + 1] as usize;
70        start..end
71    }
72
73    /// O(log n) binary-search lookup: which segment contains element `elem`?
74    ///
75    /// Returns the segment index such that `segment(idx)` contains `elem`.
76    ///
77    /// # Panics
78    /// Panics if `elem` is out of range (>= total element count).
79    pub fn segment_of(&self, elem: usize) -> usize {
80        let elem_u32 = elem as u32;
81        // Find the last offset that is <= elem. The partition point gives the
82        // first index where offsets[i] > elem, so subtract 1.
83        let pos = self.offsets.partition_point(|&o| o <= elem_u32);
84        assert!(pos > 0, "element {} is out of range", elem);
85        pos - 1
86    }
87
88    /// Iterate over all segment ranges in order.
89    pub fn iter(&self) -> impl Iterator<Item = Range<usize>> + '_ {
90        self.offsets
91            .windows(2)
92            .map(|w| (w[0] as usize)..(w[1] as usize))
93    }
94}
95
96#[cfg(test)]
97mod tests {
98    use super::*;
99
100    #[test]
101    fn test_segmentation_from_change_points() {
102        let chain_ids = ["A", "A", "A", "B", "B", "C"];
103        let seg = Segmentation::from_change_points(chain_ids.iter().copied());
104        assert_eq!(seg.count(), 3);
105        assert_eq!(seg.segment(0), 0..3);
106        assert_eq!(seg.segment(1), 3..5);
107        assert_eq!(seg.segment(2), 5..6);
108    }
109
110    #[test]
111    fn test_segmentation_from_offsets() {
112        let seg = Segmentation::from_offsets(vec![0, 3, 5, 6]);
113        assert_eq!(seg.count(), 3);
114        assert_eq!(seg.segment(0), 0..3);
115        assert_eq!(seg.segment(1), 3..5);
116        assert_eq!(seg.segment(2), 5..6);
117    }
118
119    #[test]
120    fn test_segmentation_segment_of() {
121        let seg = Segmentation::from_offsets(vec![0, 3, 5, 6]);
122        // Segment 0: elements 0,1,2
123        assert_eq!(seg.segment_of(0), 0);
124        assert_eq!(seg.segment_of(1), 0);
125        assert_eq!(seg.segment_of(2), 0);
126        // Segment 1: elements 3,4
127        assert_eq!(seg.segment_of(3), 1);
128        assert_eq!(seg.segment_of(4), 1);
129        // Segment 2: element 5
130        assert_eq!(seg.segment_of(5), 2);
131    }
132
133    #[test]
134    fn test_segmentation_round_trip() {
135        let chain_ids = ["A", "A", "A", "B", "B", "C"];
136        let seg = Segmentation::from_change_points(chain_ids.iter().copied());
137
138        // Build the expected mapping by hand
139        let expected = [0usize, 0, 0, 1, 1, 2];
140        for (elem, &exp_seg) in expected.iter().enumerate() {
141            assert_eq!(
142                seg.segment_of(elem),
143                exp_seg,
144                "elem {} should be in segment {}",
145                elem,
146                exp_seg
147            );
148        }
149    }
150
151    #[test]
152    fn test_segmentation_single_segment() {
153        let keys = ["X", "X", "X", "X"];
154        let seg = Segmentation::from_change_points(keys.iter().copied());
155        assert_eq!(seg.count(), 1);
156        assert_eq!(seg.segment(0), 0..4);
157        for elem in 0..4 {
158            assert_eq!(seg.segment_of(elem), 0);
159        }
160    }
161
162    #[test]
163    fn test_segmentation_all_different() {
164        let keys = ["A", "B", "C", "D"];
165        let seg = Segmentation::from_change_points(keys.iter().copied());
166        assert_eq!(seg.count(), 4);
167        for i in 0..4 {
168            assert_eq!(seg.segment(i), i..(i + 1));
169            assert_eq!(seg.segment_of(i), i);
170        }
171    }
172
173    #[test]
174    fn test_segmentation_iter() {
175        let seg = Segmentation::from_offsets(vec![0, 3, 5, 6]);
176        let ranges: Vec<Range<usize>> = seg.iter().collect();
177        assert_eq!(ranges, vec![0..3, 3..5, 5..6]);
178    }
179}