Skip to main content

celox_slt/
range_store.rs

1use celox_design::BitAccess;
2use serde::{Deserialize, Serialize};
3use std::collections::BTreeMap;
4use std::fmt;
5
6#[derive(Debug, Clone, PartialEq, Eq)]
7pub struct RangeStoreError {
8    message: String,
9}
10
11impl RangeStoreError {
12    fn new(message: impl Into<String>) -> Self {
13        Self {
14            message: message.into(),
15        }
16    }
17}
18
19impl fmt::Display for RangeStoreError {
20    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
21        self.message.fmt(f)
22    }
23}
24
25impl std::error::Error for RangeStoreError {}
26
27#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
28#[serde(bound(serialize = "T: Serialize", deserialize = "T: Deserialize<'de>"))]
29pub struct RangeStore<T> {
30    /// key: lsb (absolute position)
31    /// value: (expression, width, origin LSB when this data was originally placed)
32    pub ranges: BTreeMap<usize, (T, usize, usize)>,
33}
34
35pub type AlignedRangePart<'a, T, U> = (BitAccess, (&'a T, BitAccess), (&'a U, BitAccess));
36
37impl<T> RangeStore<T> {
38    fn total_width(&self) -> Result<usize, RangeStoreError> {
39        let Some((&last_lsb, (_, last_width, _))) = self.ranges.last_key_value() else {
40            return Ok(0);
41        };
42        if *last_width == 0 {
43            return Err(RangeStoreError::new(
44                "range store contains a zero-width terminal range",
45            ));
46        }
47        last_lsb
48            .checked_add(*last_width)
49            .ok_or_else(|| RangeStoreError::new("range store total width overflows usize"))
50    }
51
52    /// Align two sparse partitions in one ordered pass.
53    ///
54    /// Each returned absolute range is bounded by the next boundary from
55    /// either store. The accompanying accesses are relative to each value's
56    /// original placement, so callers can slice the values without searching
57    /// either map again.
58    pub fn aligned_parts<'a, U>(
59        &'a self,
60        other: &'a RangeStore<U>,
61    ) -> Result<Vec<AlignedRangePart<'a, T, U>>, RangeStoreError> {
62        let total_width = self.total_width()?;
63        let other_width = other.total_width()?;
64        if total_width != other_width {
65            return Err(RangeStoreError::new(format!(
66                "range store widths differ: {total_width} and {other_width}"
67            )));
68        }
69        if total_width == 0 {
70            return Ok(Vec::new());
71        }
72
73        let mut left_ranges = self.ranges.iter().peekable();
74        let mut right_ranges = other.ranges.iter().peekable();
75        let mut left = left_ranges
76            .next()
77            .ok_or_else(|| RangeStoreError::new("left range store is empty"))?;
78        let mut right = right_ranges
79            .next()
80            .ok_or_else(|| RangeStoreError::new("right range store is empty"))?;
81        if *left.0 != 0 || *right.0 != 0 {
82            return Err(RangeStoreError::new(
83                "range store does not begin at bit zero",
84            ));
85        }
86
87        let mut parts = Vec::with_capacity(self.ranges.len() + other.ranges.len());
88        let mut lsb = 0;
89        while lsb < total_width {
90            let left_boundary = left_ranges
91                .peek()
92                .map(|(next_lsb, _)| **next_lsb)
93                .unwrap_or(total_width);
94            let right_boundary = right_ranges
95                .peek()
96                .map(|(next_lsb, _)| **next_lsb)
97                .unwrap_or(total_width);
98            validate_range_boundary(left, left_boundary, total_width)?;
99            validate_range_boundary(right, right_boundary, total_width)?;
100
101            let next_lsb = left_boundary.min(right_boundary);
102            if next_lsb <= lsb {
103                return Err(RangeStoreError::new(
104                    "range store boundaries are not strictly ordered",
105                ));
106            }
107            let msb = next_lsb - 1;
108            let left_access = relative_access(left.1.2, lsb, msb)?;
109            let right_access = relative_access(right.1.2, lsb, msb)?;
110            parts.push((
111                BitAccess::new(lsb, msb),
112                (&left.1.0, left_access),
113                (&right.1.0, right_access),
114            ));
115
116            lsb = next_lsb;
117            if left_boundary == lsb && lsb < total_width {
118                left = left_ranges
119                    .next()
120                    .expect("a peeked left range boundary exists");
121            }
122            if right_boundary == lsb && lsb < total_width {
123                right = right_ranges
124                    .next()
125                    .expect("a peeked right range boundary exists");
126            }
127        }
128        Ok(parts)
129    }
130}
131
132fn validate_range_boundary<T>(
133    current: (&usize, &(T, usize, usize)),
134    next_lsb: usize,
135    total_width: usize,
136) -> Result<(), RangeStoreError> {
137    let (lsb, (_, width, _)) = current;
138    if *width == 0 {
139        return Err(RangeStoreError::new(format!(
140            "range at bit {lsb} has zero width"
141        )));
142    }
143    let end = lsb
144        .checked_add(*width)
145        .ok_or_else(|| RangeStoreError::new("range end overflows usize"))?;
146    if end != next_lsb || end > total_width {
147        return Err(RangeStoreError::new(format!(
148            "range at bit {lsb} does not end at the next boundary {next_lsb}"
149        )));
150    }
151    Ok(())
152}
153
154fn relative_access(
155    origin: usize,
156    absolute_lsb: usize,
157    absolute_msb: usize,
158) -> Result<BitAccess, RangeStoreError> {
159    let lsb = absolute_lsb
160        .checked_sub(origin)
161        .ok_or_else(|| RangeStoreError::new("range origin is above its aligned LSB"))?;
162    let msb = absolute_msb
163        .checked_sub(origin)
164        .ok_or_else(|| RangeStoreError::new("range origin is above its aligned MSB"))?;
165    Ok(BitAccess::new(lsb, msb))
166}
167
168impl<T: Clone + PartialEq + Eq> RangeStore<T> {
169    pub fn new(initial: T, width: usize) -> Self {
170        let mut ranges = BTreeMap::new();
171        if width > 0 {
172            // In initial state, absolute position 0 and origin 0 match
173            ranges.insert(0, (initial, width, 0));
174        }
175        Self { ranges }
176    }
177
178    fn validate_access(&self, access: BitAccess) -> Result<usize, RangeStoreError> {
179        let width = access
180            .msb
181            .checked_sub(access.lsb)
182            .and_then(|span| span.checked_add(1))
183            .ok_or_else(|| {
184                RangeStoreError::new(format!(
185                    "range access [{}:{}] is malformed",
186                    access.msb, access.lsb
187                ))
188            })?;
189        let total_width = self.total_width()?;
190        if total_width == 0 || access.msb >= total_width {
191            return Err(RangeStoreError::new(format!(
192                "range access [{}:{}] is outside store width {total_width}",
193                access.msb, access.lsb
194            )));
195        }
196        Ok(width)
197    }
198
199    /// Split the range at the specified bit position.
200    /// Even if split, origin_lsb (the 3rd element) is maintained.
201    pub fn split_at(&mut self, bit: usize) -> Result<(), RangeStoreError> {
202        if bit == 0 {
203            return Ok(());
204        }
205        let total_width = self.total_width()?;
206        if bit > total_width {
207            return Err(RangeStoreError::new(format!(
208                "split position {bit} is outside store width {total_width}"
209            )));
210        }
211
212        let mut split = None;
213        if let Some((&lsb, (expr, width, origin))) = self.ranges.range(..bit).next_back() {
214            if *width == 0 {
215                return Err(RangeStoreError::new(format!(
216                    "range at bit {lsb} has zero width"
217                )));
218            }
219            let msb = lsb
220                .checked_add(*width - 1)
221                .ok_or_else(|| RangeStoreError::new("range end overflows usize"))?;
222            if bit > lsb && bit <= msb {
223                // Left width: bit - lsb
224                // Right width: msb - bit + 1
225                // Both inherit the original origin
226                split = Some((lsb, bit, expr.clone(), bit - lsb, msb - bit + 1, *origin));
227            }
228        }
229
230        if let Some((lsb, bit, expr, left_w, right_w, origin)) = split {
231            self.ranges.insert(lsb, (expr.clone(), left_w, origin));
232            self.ranges.insert(bit, (expr, right_w, origin));
233        }
234        Ok(())
235    }
236
237    /// Update the specified range with a new value.
238    /// The origin_lsb of the updated range will match access.lsb of that assignment.
239    pub fn update(&mut self, access: BitAccess, value: T) -> Result<(), RangeStoreError> {
240        let width = self.validate_access(access)?;
241        let end = access
242            .msb
243            .checked_add(1)
244            .ok_or_else(|| RangeStoreError::new("updated range end overflows usize"))?;
245        self.split_at(access.lsb)?;
246        self.split_at(end)?;
247
248        self.ranges
249            .extract_if(access.lsb..=access.msb, |_, _| true)
250            .for_each(drop);
251
252        // When inserting a new range, record access.lsb as the origin
253        self.ranges.insert(access.lsb, (value, width, access.lsb));
254        Ok(())
255    }
256
257    /// Returns borrowed parts overlapping with the requested range.
258    /// relative_access will be the relative position from the origin of that expression.
259    pub fn get_parts_ref(
260        &self,
261        access: BitAccess,
262    ) -> Result<Vec<(&T, BitAccess)>, RangeStoreError> {
263        self.validate_access(access)?;
264        let mut parts = Vec::new();
265        let first_lsb = self
266            .ranges
267            .range(..=access.lsb)
268            .next_back()
269            .map(|(&lsb, _)| lsb)
270            .ok_or_else(|| RangeStoreError::new("range store does not cover access LSB"))?;
271        for (&range_lsb, (expr, range_width, origin)) in self.ranges.range(first_lsb..=access.msb) {
272            if *range_width == 0 {
273                return Err(RangeStoreError::new(format!(
274                    "range at bit {range_lsb} has zero width"
275                )));
276            }
277            let range_msb = range_lsb
278                .checked_add(*range_width - 1)
279                .ok_or_else(|| RangeStoreError::new("range end overflows usize"))?;
280
281            let overlap_lsb = range_lsb.max(access.lsb);
282            let overlap_msb = range_msb.min(access.msb);
283
284            if overlap_lsb <= overlap_msb {
285                // By subtracting origin from absolute position (overlap),
286                // calculate the correct relative index for the original data.
287                let relative_lsb = overlap_lsb.checked_sub(*origin).ok_or_else(|| {
288                    RangeStoreError::new("range origin is above its overlapping LSB")
289                })?;
290                let relative_msb = overlap_msb.checked_sub(*origin).ok_or_else(|| {
291                    RangeStoreError::new("range origin is above its overlapping MSB")
292                })?;
293                let relative_access = BitAccess::new(relative_lsb, relative_msb);
294                parts.push((expr, relative_access));
295            }
296        }
297        Ok(parts)
298    }
299
300    /// Returns owned parts overlapping with the requested range.
301    pub fn get_parts(&self, access: BitAccess) -> Result<Vec<(T, BitAccess)>, RangeStoreError> {
302        self.get_parts_ref(access).map(|parts| {
303            parts
304                .into_iter()
305                .map(|(value, access)| (value.clone(), access))
306                .collect()
307        })
308    }
309}
310
311#[cfg(test)]
312mod tests {
313    use super::*;
314
315    #[test]
316    fn rejects_malformed_and_out_of_bounds_accesses_without_panicking() {
317        let mut store = RangeStore::new(0u8, 8);
318        let original = store.clone();
319        assert!(store.update(BitAccess { lsb: 7, msb: 6 }, 1).is_err());
320        assert_eq!(store, original);
321        assert!(store.update(BitAccess::new(7, 8), 1).is_err());
322        assert_eq!(store, original);
323        assert!(store.get_parts(BitAccess::new(0, 8)).is_err());
324        assert!(store.split_at(9).is_err());
325        assert_eq!(store, original);
326    }
327
328    #[test]
329    fn checked_split_update_and_read_preserve_ranges() {
330        let mut store = RangeStore::new(0u8, 8);
331        store.update(BitAccess::new(2, 5), 1).unwrap();
332        assert_eq!(
333            store.get_parts(BitAccess::new(1, 6)).unwrap(),
334            vec![
335                (0, BitAccess::new(1, 1)),
336                (1, BitAccess::new(0, 3)),
337                (0, BitAccess::new(6, 6)),
338            ]
339        );
340    }
341
342    #[test]
343    fn reads_from_the_range_containing_the_access_lsb() {
344        let mut store = RangeStore::new(0u8, 64);
345        for bit in 0..64 {
346            store.update(BitAccess::new(bit, bit), bit as u8).unwrap();
347        }
348        assert_eq!(
349            store.get_parts(BitAccess::new(60, 62)).unwrap(),
350            vec![
351                (60, BitAccess::new(0, 0)),
352                (61, BitAccess::new(0, 0)),
353                (62, BitAccess::new(0, 0)),
354            ]
355        );
356    }
357
358    #[test]
359    fn wide_update_replaces_all_covered_sparse_ranges() {
360        let mut store = RangeStore::new(0u8, 64);
361        for bit in 0..64 {
362            store.update(BitAccess::new(bit, bit), bit as u8).unwrap();
363        }
364
365        store.update(BitAccess::new(16, 47), 99).unwrap();
366
367        assert_eq!(store.ranges.len(), 33);
368        assert_eq!(
369            store.get_parts(BitAccess::new(15, 48)).unwrap(),
370            vec![
371                (15, BitAccess::new(0, 0)),
372                (99, BitAccess::new(0, 31)),
373                (48, BitAccess::new(0, 0)),
374            ]
375        );
376    }
377
378    #[test]
379    fn aligns_two_sparse_partitions_in_boundary_order() {
380        let mut left = RangeStore::new(0u8, 8);
381        left.update(BitAccess::new(2, 5), 1).unwrap();
382        let mut right = RangeStore::new(0u8, 8);
383        right.update(BitAccess::new(4, 7), 2).unwrap();
384
385        let parts = left
386            .aligned_parts(&right)
387            .unwrap()
388            .into_iter()
389            .map(|(absolute, (left, left_access), (right, right_access))| {
390                (absolute, (*left, left_access), (*right, right_access))
391            })
392            .collect::<Vec<_>>();
393        assert_eq!(
394            parts,
395            vec![
396                (
397                    BitAccess::new(0, 1),
398                    (0, BitAccess::new(0, 1)),
399                    (0, BitAccess::new(0, 1)),
400                ),
401                (
402                    BitAccess::new(2, 3),
403                    (1, BitAccess::new(0, 1)),
404                    (0, BitAccess::new(2, 3)),
405                ),
406                (
407                    BitAccess::new(4, 5),
408                    (1, BitAccess::new(2, 3)),
409                    (2, BitAccess::new(0, 1)),
410                ),
411                (
412                    BitAccess::new(6, 7),
413                    (0, BitAccess::new(6, 7)),
414                    (2, BitAccess::new(2, 3)),
415                ),
416            ]
417        );
418    }
419
420    #[test]
421    fn aligning_partitions_rejects_different_widths() {
422        let left = RangeStore::new(0u8, 8);
423        let right = RangeStore::new(0u8, 9);
424
425        assert!(left.aligned_parts(&right).is_err());
426    }
427}