Skip to main content

sbi_spec/binary/
hart_mask.rs

1use super::{
2    mask_commons::{MaskError, has_bit, valid_bit},
3    sbi_ret::SbiRegister,
4};
5
6/// Hart mask structure in SBI function calls.
7#[repr(C)]
8#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)]
9pub struct HartMask<T = usize> {
10    hart_mask: T,
11    hart_mask_base: T,
12}
13
14impl<T: SbiRegister> HartMask<T> {
15    /// Special value to ignore the `mask`, and consider all `bit`s as set.
16    pub const IGNORE_MASK: T = T::FULL_MASK;
17
18    /// Construct a [HartMask] from mask value and base hart id.
19    #[inline]
20    pub const fn from_mask_base(hart_mask: T, hart_mask_base: T) -> Self {
21        Self {
22            hart_mask,
23            hart_mask_base,
24        }
25    }
26
27    /// Construct a [HartMask] that selects all available harts on the current environment.
28    ///
29    /// According to the RISC-V SBI Specification, `hart_mask_base` can be set to `-1` (i.e. `usize::MAX`)
30    /// to indicate that `hart_mask` shall be ignored and all available harts must be considered.
31    /// In case of this function in the `sbi-spec` crate, we fill in `usize::MAX` in `hart_mask_base`
32    /// parameter to match the RISC-V SBI standard, while choosing 0 as the ignored `hart_mask` value.
33    #[inline]
34    pub const fn all() -> Self {
35        Self {
36            hart_mask: T::ZERO,
37            hart_mask_base: T::FULL_MASK,
38        }
39    }
40
41    /// Gets the special value for ignoring the `mask` parameter.
42    #[inline]
43    pub const fn ignore_mask(&self) -> T {
44        Self::IGNORE_MASK
45    }
46
47    /// Returns `mask` and `base` parameters from the [HartMask].
48    #[inline]
49    pub const fn into_inner(self) -> (T, T) {
50        (self.hart_mask, self.hart_mask_base)
51    }
52}
53
54// FIXME: implement for T: SbiRegister once we can implement this using const traits.
55// Ref: https://rust-lang.github.io/rust-project-goals/2024h2/const-traits.html
56impl HartMask<usize> {
57    /// Returns whether the [HartMask] contains the provided `hart_id`.
58    #[inline]
59    pub const fn has_bit(self, hart_id: usize) -> bool {
60        has_bit(
61            self.hart_mask,
62            self.hart_mask_base,
63            Self::IGNORE_MASK,
64            hart_id,
65        )
66    }
67
68    /// Insert a hart id into this [HartMask].
69    ///
70    /// Returns error when `hart_id` is invalid.
71    #[inline]
72    pub const fn insert(&mut self, hart_id: usize) -> Result<(), MaskError> {
73        if self.hart_mask_base == Self::IGNORE_MASK {
74            Ok(())
75        } else if valid_bit(self.hart_mask_base, hart_id) {
76            self.hart_mask |= 1usize << (hart_id - self.hart_mask_base);
77            Ok(())
78        } else {
79            Err(MaskError::InvalidBit)
80        }
81    }
82
83    /// Remove a hart id from this [HartMask].
84    ///
85    /// Returns error when `hart_id` is invalid, or it has been ignored.
86    #[inline]
87    pub const fn remove(&mut self, hart_id: usize) -> Result<(), MaskError> {
88        if self.hart_mask_base == Self::IGNORE_MASK {
89            Err(MaskError::Ignored)
90        } else if valid_bit(self.hart_mask_base, hart_id) {
91            self.hart_mask &= !(1usize << (hart_id - self.hart_mask_base));
92            Ok(())
93        } else {
94            Err(MaskError::InvalidBit)
95        }
96    }
97
98    /// Returns [HartIds] of self.
99    #[inline]
100    pub const fn iter(&self) -> HartIds {
101        HartIds {
102            inner: match self.hart_mask_base {
103                Self::IGNORE_MASK => UnvisitedMask::Range(0, usize::MAX),
104                _ => UnvisitedMask::MaskBase(self.hart_mask, self.hart_mask_base),
105            },
106        }
107    }
108}
109
110impl IntoIterator for HartMask {
111    type Item = usize;
112
113    type IntoIter = HartIds;
114
115    #[inline]
116    fn into_iter(self) -> Self::IntoIter {
117        self.iter()
118    }
119}
120
121/// Iterator structure for `HartMask`.
122///
123/// It will iterate hart id from low to high.
124#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)]
125pub struct HartIds {
126    inner: UnvisitedMask,
127}
128
129#[derive(Debug, Copy, Clone, Eq, PartialEq, Hash)]
130enum UnvisitedMask {
131    MaskBase(usize, usize),
132    Range(usize, usize),
133}
134
135impl Iterator for HartIds {
136    type Item = usize;
137
138    #[inline]
139    fn next(&mut self) -> Option<Self::Item> {
140        match &mut self.inner {
141            UnvisitedMask::MaskBase(0, _base) => None,
142            UnvisitedMask::MaskBase(unvisited_mask, base) => {
143                let low_bit = unvisited_mask.trailing_zeros();
144                let hart_id = usize::try_from(low_bit).unwrap() + *base;
145                *unvisited_mask &= !(1usize << low_bit);
146                Some(hart_id)
147            }
148            UnvisitedMask::Range(start, end) => {
149                assert!(start <= end);
150                if *start < *end {
151                    let ans = *start;
152                    *start += 1;
153                    Some(ans)
154                } else {
155                    None
156                }
157            }
158        }
159    }
160
161    #[inline]
162    fn size_hint(&self) -> (usize, Option<usize>) {
163        match self.inner {
164            UnvisitedMask::MaskBase(unvisited_mask, _base) => {
165                let exact_popcnt = usize::try_from(unvisited_mask.count_ones()).unwrap();
166                (exact_popcnt, Some(exact_popcnt))
167            }
168            UnvisitedMask::Range(start, end) => {
169                assert!(start <= end);
170                let exact_num_harts = end - start;
171                (exact_num_harts, Some(exact_num_harts))
172            }
173        }
174    }
175
176    #[inline]
177    fn count(self) -> usize {
178        self.size_hint().0
179    }
180
181    #[inline]
182    fn last(mut self) -> Option<Self::Item> {
183        self.next_back()
184    }
185
186    #[inline]
187    fn min(mut self) -> Option<Self::Item> {
188        self.next()
189    }
190
191    #[inline]
192    fn max(mut self) -> Option<Self::Item> {
193        self.next_back()
194    }
195
196    #[inline]
197    fn is_sorted(self) -> bool {
198        true
199    }
200
201    // TODO: implement fn advance_by once it's stabilized: https://github.com/rust-lang/rust/issues/77404
202    // #[inline]
203    // fn advance_by(&mut self, n: usize) -> Result<(), core::num::NonZero<usize>> { ... }
204}
205
206impl DoubleEndedIterator for HartIds {
207    #[inline]
208    fn next_back(&mut self) -> Option<Self::Item> {
209        match &mut self.inner {
210            UnvisitedMask::MaskBase(0, _base) => None,
211            UnvisitedMask::MaskBase(unvisited_mask, base) => {
212                let high_bit = unvisited_mask.leading_zeros();
213                let hart_id = usize::try_from(usize::BITS - high_bit - 1).unwrap() + *base;
214                *unvisited_mask &= !(1usize << (usize::BITS - high_bit - 1));
215                Some(hart_id)
216            }
217            UnvisitedMask::Range(start, end) => {
218                assert!(start <= end);
219                if *start < *end {
220                    let ans = *end;
221                    *end -= 1;
222                    Some(ans)
223                } else {
224                    None
225                }
226            }
227        }
228    }
229
230    // TODO: implement advance_back_by once stabilized.
231    // #[inline]
232    // fn advance_back_by(&mut self, n: usize) -> Result<(), core::num::NonZero<usize>> { ... }
233}
234
235impl ExactSizeIterator for HartIds {}
236
237impl core::iter::FusedIterator for HartIds {}
238
239#[cfg(test)]
240mod tests {
241    use super::*;
242
243    #[test]
244    fn rustsbi_hart_mask() {
245        let mask = HartMask::from_mask_base(0b1, 400);
246        assert!(!mask.has_bit(0));
247        assert!(mask.has_bit(400));
248        assert!(!mask.has_bit(401));
249        let mask = HartMask::from_mask_base(0b110, 500);
250        assert!(!mask.has_bit(0));
251        assert!(!mask.has_bit(500));
252        assert!(mask.has_bit(501));
253        assert!(mask.has_bit(502));
254        assert!(!mask.has_bit(500 + (usize::BITS as usize)));
255        let max_bit = 1 << (usize::BITS - 1);
256        let mask = HartMask::from_mask_base(max_bit, 600);
257        assert!(mask.has_bit(600 + (usize::BITS as usize) - 1));
258        assert!(!mask.has_bit(600 + (usize::BITS as usize)));
259        let mask = HartMask::from_mask_base(0b11, usize::MAX - 1);
260        assert!(!mask.has_bit(usize::MAX - 2));
261        assert!(mask.has_bit(usize::MAX - 1));
262        assert!(mask.has_bit(usize::MAX));
263        assert!(!mask.has_bit(0));
264        // hart_mask_base == usize::MAX is special, it means hart_mask should be ignored
265        // and this hart mask contains all harts available
266        let mask = HartMask::from_mask_base(0, usize::MAX);
267        for i in 0..5 {
268            assert!(mask.has_bit(i));
269        }
270        assert!(mask.has_bit(usize::MAX));
271
272        let mut mask = HartMask::from_mask_base(0, 1);
273        assert!(!mask.has_bit(1));
274        assert!(mask.insert(1).is_ok());
275        assert!(mask.has_bit(1));
276        assert!(mask.remove(1).is_ok());
277        assert!(!mask.has_bit(1));
278    }
279
280    #[test]
281    fn rustsbi_hart_ids_iterator() {
282        let mask = HartMask::from_mask_base(0b101011, 1);
283        // Test the `next` method of `HartIds` structure.
284        let mut hart_ids = mask.iter();
285        assert_eq!(hart_ids.next(), Some(1));
286        assert_eq!(hart_ids.next(), Some(2));
287        assert_eq!(hart_ids.next(), Some(4));
288        assert_eq!(hart_ids.next(), Some(6));
289        assert_eq!(hart_ids.next(), None);
290        // `HartIds` structures are fused, meaning they return `None` forever once iteration finished.
291        assert_eq!(hart_ids.next(), None);
292
293        // Test `for` loop on mask (`HartMask`) as `IntoIterator`.
294        let mut ans = [0; 4];
295        let mut idx = 0;
296        for hart_id in mask {
297            ans[idx] = hart_id;
298            idx += 1;
299        }
300        assert_eq!(ans, [1, 2, 4, 6]);
301
302        // Test `Iterator` methods on `HartIds`.
303        let mut hart_ids = mask.iter();
304        assert_eq!(hart_ids.size_hint(), (4, Some(4)));
305        let _ = hart_ids.next();
306        assert_eq!(hart_ids.size_hint(), (3, Some(3)));
307        let _ = hart_ids.next();
308        let _ = hart_ids.next();
309        assert_eq!(hart_ids.size_hint(), (1, Some(1)));
310        let _ = hart_ids.next();
311        assert_eq!(hart_ids.size_hint(), (0, Some(0)));
312        let _ = hart_ids.next();
313        assert_eq!(hart_ids.size_hint(), (0, Some(0)));
314
315        let mut hart_ids = mask.iter();
316        assert_eq!(hart_ids.count(), 4);
317        let _ = hart_ids.next();
318        assert_eq!(hart_ids.count(), 3);
319        let _ = hart_ids.next();
320        let _ = hart_ids.next();
321        let _ = hart_ids.next();
322        assert_eq!(hart_ids.count(), 0);
323        let _ = hart_ids.next();
324        assert_eq!(hart_ids.count(), 0);
325
326        let hart_ids = mask.iter();
327        assert_eq!(hart_ids.last(), Some(6));
328
329        let mut hart_ids = mask.iter();
330        assert_eq!(hart_ids.nth(2), Some(4));
331        let mut hart_ids = mask.iter();
332        assert_eq!(hart_ids.nth(0), Some(1));
333
334        let mut iter = mask.iter().step_by(2);
335        assert_eq!(iter.next(), Some(1));
336        assert_eq!(iter.next(), Some(4));
337        assert_eq!(iter.next(), None);
338
339        let mask_2 = HartMask::from_mask_base(0b1001101, 64);
340        let mut iter = mask.iter().chain(mask_2);
341        assert_eq!(iter.next(), Some(1));
342        assert_eq!(iter.next(), Some(2));
343        assert_eq!(iter.next(), Some(4));
344        assert_eq!(iter.next(), Some(6));
345        assert_eq!(iter.next(), Some(64));
346        assert_eq!(iter.next(), Some(66));
347        assert_eq!(iter.next(), Some(67));
348        assert_eq!(iter.next(), Some(70));
349        assert_eq!(iter.next(), None);
350
351        let mut iter = mask.iter().zip(mask_2);
352        assert_eq!(iter.next(), Some((1, 64)));
353        assert_eq!(iter.next(), Some((2, 66)));
354        assert_eq!(iter.next(), Some((4, 67)));
355        assert_eq!(iter.next(), Some((6, 70)));
356        assert_eq!(iter.next(), None);
357
358        fn to_plic_context_id(hart_id_machine: usize) -> usize {
359            hart_id_machine * 2
360        }
361        let mut iter = mask.iter().map(to_plic_context_id);
362        assert_eq!(iter.next(), Some(2));
363        assert_eq!(iter.next(), Some(4));
364        assert_eq!(iter.next(), Some(8));
365        assert_eq!(iter.next(), Some(12));
366        assert_eq!(iter.next(), None);
367
368        let mut channel_received = [0; 4];
369        let mut idx = 0;
370        let mut channel_send = |hart_id| {
371            channel_received[idx] = hart_id;
372            idx += 1;
373        };
374        mask.iter().for_each(|value| channel_send(value));
375        assert_eq!(channel_received, [1, 2, 4, 6]);
376
377        let is_in_cluster_1 = |hart_id: &usize| *hart_id >= 4 && *hart_id < 7;
378        let mut iter = mask.iter().filter(is_in_cluster_1);
379        assert_eq!(iter.next(), Some(4));
380        assert_eq!(iter.next(), Some(6));
381        assert_eq!(iter.next(), None);
382
383        let if_in_cluster_1_get_plic_context_id = |hart_id: usize| {
384            if hart_id >= 4 && hart_id < 7 {
385                Some(hart_id * 2)
386            } else {
387                None
388            }
389        };
390        let mut iter = mask.iter().filter_map(if_in_cluster_1_get_plic_context_id);
391        assert_eq!(iter.next(), Some(8));
392        assert_eq!(iter.next(), Some(12));
393        assert_eq!(iter.next(), None);
394
395        let mut iter = mask.iter().enumerate();
396        assert_eq!(iter.next(), Some((0, 1)));
397        assert_eq!(iter.next(), Some((1, 2)));
398        assert_eq!(iter.next(), Some((2, 4)));
399        assert_eq!(iter.next(), Some((3, 6)));
400        assert_eq!(iter.next(), None);
401        let mut ans = [(0, 0); 4];
402        let mut idx = 0;
403        for (i, hart_id) in mask.iter().enumerate() {
404            ans[idx] = (i, hart_id);
405            idx += 1;
406        }
407        assert_eq!(ans, [(0, 1), (1, 2), (2, 4), (3, 6)]);
408
409        let mut iter = mask.iter().peekable();
410        assert_eq!(iter.peek(), Some(&1));
411        assert_eq!(iter.next(), Some(1));
412        assert_eq!(iter.peek(), Some(&2));
413        assert_eq!(iter.next(), Some(2));
414        assert_eq!(iter.peek(), Some(&4));
415        assert_eq!(iter.next(), Some(4));
416        assert_eq!(iter.peek(), Some(&6));
417        assert_eq!(iter.next(), Some(6));
418        assert_eq!(iter.peek(), None);
419        assert_eq!(iter.next(), None);
420
421        // TODO: other iterator tests.
422
423        assert!(mask.iter().is_sorted());
424        assert!(mask.iter().is_sorted_by(|a, b| a <= b));
425
426        // Reverse iterator as `DoubleEndedIterator`.
427        let mut iter = mask.iter().rev();
428        assert_eq!(iter.next(), Some(6));
429        assert_eq!(iter.next(), Some(4));
430        assert_eq!(iter.next(), Some(2));
431        assert_eq!(iter.next(), Some(1));
432        assert_eq!(iter.next(), None);
433
434        // Special iterator values.
435        let nothing = HartMask::from_mask_base(0, 1000);
436        assert!(nothing.iter().eq([]));
437
438        let all_mask_bits_set = HartMask::from_mask_base(usize::MAX, 1000);
439        let range = 1000..(1000 + usize::BITS as usize);
440        assert!(all_mask_bits_set.iter().eq(range));
441
442        let all_harts = HartMask::all();
443        let mut iter = all_harts.iter();
444        assert_eq!(iter.size_hint(), (usize::MAX, Some(usize::MAX)));
445        // Don't use `Iterator::eq` here; it would literally run `Iterator::try_for_each` from 0 to usize::MAX
446        // which will cost us forever to run the test.
447        assert_eq!(iter.next(), Some(0));
448        assert_eq!(iter.size_hint(), (usize::MAX - 1, Some(usize::MAX - 1)));
449        assert_eq!(iter.next(), Some(1));
450        assert_eq!(iter.next(), Some(2));
451        // skip 500 elements
452        let _ = iter.nth(500 - 1);
453        assert_eq!(iter.next(), Some(503));
454        assert_eq!(iter.size_hint(), (usize::MAX - 504, Some(usize::MAX - 504)));
455        assert_eq!(iter.next_back(), Some(usize::MAX));
456        assert_eq!(iter.next_back(), Some(usize::MAX - 1));
457        assert_eq!(iter.size_hint(), (usize::MAX - 506, Some(usize::MAX - 506)));
458
459        // A common usage of `HartMask::all`, we assume that this platform filters out hart 0..=3.
460        let environment_available_hart_ids = 4..128;
461        // `hart_mask_iter` contains 64..=usize::MAX.
462        let hart_mask_iter = all_harts.iter().skip(64);
463        let filtered_iter = environment_available_hart_ids.filter(|&x| {
464            hart_mask_iter
465                .clone()
466                .find(|&y| y >= x)
467                .map_or(false, |y| y == x)
468        });
469        assert!(filtered_iter.eq(64..128));
470
471        // The following operations should have O(1) complexity.
472        let all_harts = HartMask::all();
473        assert_eq!(all_harts.iter().count(), usize::MAX);
474        assert_eq!(all_harts.iter().last(), Some(usize::MAX));
475        assert_eq!(all_harts.iter().min(), Some(0));
476        assert_eq!(all_harts.iter().max(), Some(usize::MAX));
477        assert!(all_harts.iter().is_sorted());
478
479        let partial_all_harts = {
480            let mut ans = HartMask::all().iter();
481            let _ = ans.nth(65536 - 1);
482            let _ = ans.nth_back(4096 - 1);
483            ans
484        };
485        assert_eq!(partial_all_harts.clone().count(), usize::MAX - 65536 - 4096);
486        assert_eq!(partial_all_harts.clone().last(), Some(usize::MAX - 4096));
487        assert_eq!(partial_all_harts.clone().min(), Some(65536));
488        assert_eq!(partial_all_harts.clone().max(), Some(usize::MAX - 4096));
489        assert!(partial_all_harts.is_sorted());
490
491        let nothing = HartMask::from_mask_base(0, 1000);
492        assert_eq!(nothing.iter().count(), 0);
493        assert_eq!(nothing.iter().last(), None);
494        assert_eq!(nothing.iter().min(), None);
495        assert_eq!(nothing.iter().max(), None);
496        assert!(nothing.iter().is_sorted());
497
498        let mask = HartMask::from_mask_base(0b101011, 1);
499        assert_eq!(mask.iter().count(), 4);
500        assert_eq!(mask.iter().last(), Some(6));
501        assert_eq!(mask.iter().min(), Some(1));
502        assert_eq!(mask.iter().max(), Some(6));
503        assert!(mask.iter().is_sorted());
504
505        let all_mask_bits_set = HartMask::from_mask_base(usize::MAX, 1000);
506        let last = 1000 + usize::BITS as usize - 1;
507        assert_eq!(all_mask_bits_set.iter().count(), usize::BITS as usize);
508        assert_eq!(all_mask_bits_set.iter().last(), Some(last));
509        assert_eq!(all_mask_bits_set.iter().min(), Some(1000));
510        assert_eq!(all_mask_bits_set.iter().max(), Some(last));
511        assert!(all_mask_bits_set.iter().is_sorted());
512    }
513
514    #[test]
515    fn rustsbi_hart_mask_non_usize() {
516        assert_eq!(HartMask::<i32>::IGNORE_MASK, -1);
517        assert_eq!(HartMask::<i64>::IGNORE_MASK, -1);
518        assert_eq!(HartMask::<i128>::IGNORE_MASK, -1);
519        assert_eq!(HartMask::<u32>::IGNORE_MASK, u32::MAX);
520        assert_eq!(HartMask::<u64>::IGNORE_MASK, u64::MAX);
521        assert_eq!(HartMask::<u128>::IGNORE_MASK, u128::MAX);
522
523        assert_eq!(HartMask::<i32>::all(), HartMask::from_mask_base(0, -1));
524    }
525}