1use super::{
2 mask_commons::{MaskError, has_bit, valid_bit},
3 sbi_ret::SbiRegister,
4};
5
6#[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 pub const IGNORE_MASK: T = T::FULL_MASK;
17
18 #[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 #[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 #[inline]
43 pub const fn ignore_mask(&self) -> T {
44 Self::IGNORE_MASK
45 }
46
47 #[inline]
49 pub const fn into_inner(self) -> (T, T) {
50 (self.hart_mask, self.hart_mask_base)
51 }
52}
53
54impl HartMask<usize> {
57 #[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 #[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 #[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 #[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#[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 }
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 }
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 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 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 assert_eq!(hart_ids.next(), None);
292
293 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 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 assert!(mask.iter().is_sorted());
424 assert!(mask.iter().is_sorted_by(|a, b| a <= b));
425
426 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 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 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 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 let environment_available_hart_ids = 4..128;
461 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 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}