1use core::{
4 cell::UnsafeCell,
5 fmt,
6 marker::PhantomData,
7 ops::{Deref, DerefMut},
8 sync::atomic::{AtomicUsize, Ordering},
9};
10
11use ax_kernel_guard::BaseGuard;
12#[cfg(feature = "lockdep")]
13use ax_kernel_guard::IrqSave;
14
15#[cfg(feature = "lockdep")]
16type LockdepAcquire = crate::lockdep::Lockdep;
17
18#[cfg(not(feature = "lockdep"))]
19#[derive(Clone, Copy)]
20struct LockdepAcquire;
21
22#[cfg(not(feature = "lockdep"))]
23impl LockdepAcquire {
24 #[inline(always)]
25 #[track_caller]
26 fn prepare<G: BaseGuard, T: ?Sized>(_lock: &BaseSpinRwLock<G, T>, _is_try: bool) -> Self {
27 Self
28 }
29
30 #[inline(always)]
31 fn finish(&self, _acquired: bool) {}
32}
33
34const READER: usize = 1;
35const WRITER: usize = 1 << (usize::BITS - 1);
36const MAX_READER: usize = 1 << (usize::BITS - 2);
37
38pub struct BaseSpinRwLock<G: BaseGuard, T: ?Sized> {
45 _phantom: PhantomData<G>,
46 state: AtomicUsize,
47 #[cfg(feature = "lockdep")]
48 lockdep: crate::lockdep::LockdepMap,
49 data: UnsafeCell<T>,
50}
51
52pub struct BaseSpinRwLockReadGuard<'a, G: BaseGuard, T: ?Sized + 'a> {
54 _phantom: &'a PhantomData<G>,
55 guard_state: G::State,
56 #[cfg(feature = "lockdep")]
57 lock_addr: usize,
58 data: *const T,
59 state: &'a AtomicUsize,
60}
61
62pub struct BaseSpinRwLockWriteGuard<'a, G: BaseGuard, T: ?Sized + 'a> {
64 _phantom: &'a PhantomData<G>,
65 guard_state: G::State,
66 #[cfg(feature = "lockdep")]
67 lock_addr: usize,
68 data: *mut T,
69 state: &'a AtomicUsize,
70}
71
72unsafe impl<G: BaseGuard, T: ?Sized + Send> Send for BaseSpinRwLock<G, T> {}
73unsafe impl<G: BaseGuard, T: ?Sized + Send + Sync> Sync for BaseSpinRwLock<G, T> {}
74
75impl<G: BaseGuard, T> BaseSpinRwLock<G, T> {
76 #[inline(always)]
78 #[track_caller]
79 pub const fn new(data: T) -> Self {
80 Self {
81 _phantom: PhantomData,
82 state: AtomicUsize::new(0),
83 #[cfg(feature = "lockdep")]
84 lockdep: crate::lockdep::LockdepMap::new(),
85 data: UnsafeCell::new(data),
86 }
87 }
88
89 #[inline(always)]
91 pub fn into_inner(self) -> T {
92 let BaseSpinRwLock { data, .. } = self;
93 data.into_inner()
94 }
95}
96
97impl<G: BaseGuard, T: ?Sized> BaseSpinRwLock<G, T> {
98 #[cfg(feature = "lockdep")]
99 #[inline(always)]
100 pub(crate) fn lockdep_map(&self) -> &crate::lockdep::LockdepMap {
101 &self.lockdep
102 }
103
104 #[cfg(feature = "lockdep")]
105 #[inline(always)]
106 fn lock_addr(&self) -> usize {
107 self as *const _ as *const () as usize
108 }
109
110 #[inline(always)]
111 #[track_caller]
112 fn prepare_lockdep(&self, is_try: bool, track_task_lock: bool) -> LockdepAcquire {
113 #[cfg(not(feature = "lockdep"))]
114 let _ = track_task_lock;
115
116 #[cfg(feature = "lockdep")]
117 {
118 LockdepAcquire::prepare_map::<G>(
119 self.lockdep_map(),
120 "spin rwlock",
121 "spin-rwlock",
122 self.lock_addr(),
123 is_try,
124 crate::lockdep::DEFAULT_LOCK_SUBCLASS,
125 track_task_lock,
126 )
127 }
128
129 #[cfg(not(feature = "lockdep"))]
130 {
131 LockdepAcquire::prepare(self, is_try)
132 }
133 }
134
135 #[inline(always)]
136 fn finish_lockdep(lockdep: LockdepAcquire, acquired: bool) {
137 #[cfg(feature = "lockdep")]
138 {
139 let _lockdep_irq_guard = IrqSave::new();
140 lockdep.finish(acquired);
141 }
142
143 #[cfg(not(feature = "lockdep"))]
144 {
145 lockdep.finish(acquired);
146 }
147 }
148
149 #[inline(always)]
150 fn try_acquire_read(&self) -> bool {
151 let old = self.state.fetch_add(READER, Ordering::Acquire);
152 if old & (WRITER | MAX_READER) == 0 {
153 true
154 } else {
155 self.state.fetch_sub(READER, Ordering::Release);
156 false
157 }
158 }
159
160 #[inline(always)]
161 fn try_acquire_write(&self) -> bool {
162 self.state
163 .compare_exchange(0, WRITER, Ordering::Acquire, Ordering::Relaxed)
164 .is_ok()
165 }
166
167 #[inline(always)]
169 #[track_caller]
170 pub fn read(&self) -> BaseSpinRwLockReadGuard<'_, G, T> {
171 let guard_state = G::acquire();
172 let lockdep = self.prepare_lockdep(false, false);
173 while !self.try_acquire_read() {
174 while self.is_write_locked() {
175 core::hint::spin_loop();
176 }
177 }
178 Self::finish_lockdep(lockdep, true);
179 BaseSpinRwLockReadGuard {
180 _phantom: &PhantomData,
181 guard_state,
182 #[cfg(feature = "lockdep")]
183 lock_addr: lockdep.lock_addr(),
184 data: self.data.get(),
185 state: &self.state,
186 }
187 }
188
189 #[inline(always)]
191 #[track_caller]
192 pub fn write(&self) -> BaseSpinRwLockWriteGuard<'_, G, T> {
193 let guard_state = G::acquire();
194 let lockdep = self.prepare_lockdep(false, true);
195 while !self.try_acquire_write() {
196 while self.state.load(Ordering::Acquire) != 0 {
197 core::hint::spin_loop();
198 }
199 }
200 Self::finish_lockdep(lockdep, true);
201 BaseSpinRwLockWriteGuard {
202 _phantom: &PhantomData,
203 guard_state,
204 #[cfg(feature = "lockdep")]
205 lock_addr: lockdep.lock_addr(),
206 data: self.data.get(),
207 state: &self.state,
208 }
209 }
210
211 #[inline(always)]
213 #[track_caller]
214 pub fn try_read(&self) -> Option<BaseSpinRwLockReadGuard<'_, G, T>> {
215 let guard_state = G::acquire();
216 let lockdep = self.prepare_lockdep(true, false);
217 let acquired = self.try_acquire_read();
218 Self::finish_lockdep(lockdep, acquired);
219
220 if acquired {
221 Some(BaseSpinRwLockReadGuard {
222 _phantom: &PhantomData,
223 guard_state,
224 #[cfg(feature = "lockdep")]
225 lock_addr: lockdep.lock_addr(),
226 data: self.data.get(),
227 state: &self.state,
228 })
229 } else {
230 G::release(guard_state);
231 None
232 }
233 }
234
235 #[inline(always)]
237 #[track_caller]
238 pub fn try_write(&self) -> Option<BaseSpinRwLockWriteGuard<'_, G, T>> {
239 let guard_state = G::acquire();
240 let lockdep = self.prepare_lockdep(true, true);
241 let acquired = self.try_acquire_write();
242 Self::finish_lockdep(lockdep, acquired);
243
244 if acquired {
245 Some(BaseSpinRwLockWriteGuard {
246 _phantom: &PhantomData,
247 guard_state,
248 #[cfg(feature = "lockdep")]
249 lock_addr: lockdep.lock_addr(),
250 data: self.data.get(),
251 state: &self.state,
252 })
253 } else {
254 G::release(guard_state);
255 None
256 }
257 }
258
259 #[inline(always)]
261 pub fn is_write_locked(&self) -> bool {
262 self.state.load(Ordering::Acquire) & WRITER != 0
263 }
264
265 #[inline(always)]
270 pub fn reader_count(&self) -> usize {
271 self.state.load(Ordering::Relaxed) & !(WRITER | MAX_READER)
273 }
274
275 #[inline(always)]
280 pub fn writer_count(&self) -> usize {
281 usize::from(self.state.load(Ordering::Relaxed) & WRITER != 0)
283 }
284
285 #[inline(always)]
295 pub unsafe fn force_read_decrement(&self) {
296 let mut state = self.state.load(Ordering::Acquire);
297 loop {
298 let readers = state & !(WRITER | MAX_READER);
299 if readers == 0 {
300 return;
301 }
302
303 match self.state.compare_exchange_weak(
304 state,
305 state - READER,
306 Ordering::Release,
307 Ordering::Relaxed,
308 ) {
309 Ok(_) => {
310 #[cfg(feature = "lockdep")]
311 {
312 let _lockdep_irq_guard = IrqSave::new();
313 crate::lockdep::release_trace_only::<G>("spin-rwlock", self.lock_addr());
314 }
315 return;
316 }
317 Err(observed) => state = observed,
318 }
319 }
320 }
321
322 #[inline(always)]
329 pub unsafe fn force_write_unlock(&self) {
330 debug_assert_eq!(self.state.load(Ordering::Relaxed), WRITER);
331 #[cfg(feature = "lockdep")]
332 {
333 let _lockdep_irq_guard = IrqSave::new();
334 crate::lockdep::release_kind::<G>("spin-rwlock", self.lock_addr());
335 }
336 self.state.fetch_and(!WRITER, Ordering::Release);
337 }
338
339 #[inline(always)]
341 pub fn get_mut(&mut self) -> &mut T {
342 self.data.get_mut()
343 }
344}
345
346impl<G: BaseGuard, T: Default> Default for BaseSpinRwLock<G, T> {
347 #[inline(always)]
348 fn default() -> Self {
349 Self::new(Default::default())
350 }
351}
352
353impl<G: BaseGuard, T> From<T> for BaseSpinRwLock<G, T> {
354 #[inline(always)]
355 fn from(value: T) -> Self {
356 Self::new(value)
357 }
358}
359
360impl<G: BaseGuard, T: ?Sized + fmt::Debug> fmt::Debug for BaseSpinRwLock<G, T> {
361 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
362 match self.try_read() {
363 Some(guard) => f
364 .debug_struct("SpinRwLock")
365 .field("data", &&*guard)
366 .finish(),
367 None => write!(f, "SpinRwLock {{ <locked> }}"),
368 }
369 }
370}
371
372impl<G: BaseGuard, T: ?Sized> Deref for BaseSpinRwLockReadGuard<'_, G, T> {
373 type Target = T;
374
375 #[inline(always)]
376 fn deref(&self) -> &T {
377 unsafe { &*self.data }
378 }
379}
380
381impl<G: BaseGuard, T: ?Sized + fmt::Debug> fmt::Debug for BaseSpinRwLockReadGuard<'_, G, T> {
382 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
383 fmt::Debug::fmt(&**self, f)
384 }
385}
386
387impl<G: BaseGuard, T: ?Sized> Drop for BaseSpinRwLockReadGuard<'_, G, T> {
388 #[inline(always)]
389 fn drop(&mut self) {
390 #[cfg(feature = "lockdep")]
391 {
392 let _lockdep_irq_guard = IrqSave::new();
393 crate::lockdep::release_trace_only::<G>("spin-rwlock", self.lock_addr);
394 }
395 self.state.fetch_sub(READER, Ordering::Release);
396 G::release(self.guard_state);
397 }
398}
399
400impl<G: BaseGuard, T: ?Sized> Deref for BaseSpinRwLockWriteGuard<'_, G, T> {
401 type Target = T;
402
403 #[inline(always)]
404 fn deref(&self) -> &T {
405 unsafe { &*self.data }
406 }
407}
408
409impl<G: BaseGuard, T: ?Sized> DerefMut for BaseSpinRwLockWriteGuard<'_, G, T> {
410 #[inline(always)]
411 fn deref_mut(&mut self) -> &mut Self::Target {
412 unsafe { &mut *self.data }
413 }
414}
415
416impl<G: BaseGuard, T: ?Sized + fmt::Debug> fmt::Debug for BaseSpinRwLockWriteGuard<'_, G, T> {
417 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
418 fmt::Debug::fmt(&**self, f)
419 }
420}
421
422impl<G: BaseGuard, T: ?Sized> Drop for BaseSpinRwLockWriteGuard<'_, G, T> {
423 #[inline(always)]
424 fn drop(&mut self) {
425 #[cfg(feature = "lockdep")]
426 {
427 let _lockdep_irq_guard = IrqSave::new();
428 crate::lockdep::release_kind::<G>("spin-rwlock", self.lock_addr);
429 }
430 self.state.fetch_and(!WRITER, Ordering::Release);
431 G::release(self.guard_state);
432 }
433}
434
435#[cfg(test)]
436mod tests {
437 use std::{
438 sync::{
439 Arc,
440 atomic::{AtomicUsize, Ordering},
441 },
442 thread,
443 };
444
445 type RwLock<T> = crate::SpinRawRwLock<T>;
446
447 #[test]
448 fn readers_can_share() {
449 let lock = RwLock::new(7);
450 let first = lock.read();
451 let second = lock.try_read().expect("second reader should enter");
452
453 assert_eq!(*first, 7);
454 assert_eq!(*second, 7);
455 assert!(lock.try_write().is_none());
456 }
457
458 #[test]
459 fn writer_excludes_readers_and_writers() {
460 let lock = RwLock::new(1);
461 let mut writer = lock.write();
462 *writer = 2;
463
464 assert!(lock.try_read().is_none());
465 assert!(lock.try_write().is_none());
466 drop(writer);
467
468 assert_eq!(*lock.read(), 2);
469 }
470
471 #[test]
472 fn try_write_waits_for_all_readers() {
473 let lock = RwLock::new(());
474 let first = lock.read();
475 let second = lock.read();
476
477 assert!(lock.try_write().is_none());
478 drop(first);
479 assert!(lock.try_write().is_none());
480 drop(second);
481 assert!(lock.try_write().is_some());
482 }
483
484 #[test]
485 fn force_read_decrement_releases_leaked_reader() {
486 let lock = RwLock::new(());
487 let guard = lock.read();
488 core::mem::forget(guard);
489
490 assert_eq!(lock.reader_count(), 1);
491 assert!(lock.try_write().is_none());
492
493 unsafe { lock.force_read_decrement() };
494 assert_eq!(lock.reader_count(), 0);
495 assert!(lock.try_write().is_some());
496 }
497
498 #[test]
499 fn force_read_decrement_without_reader_does_not_poison_state() {
500 let lock = RwLock::new(());
501 let guard = lock.read();
502 core::mem::forget(guard);
503
504 unsafe { lock.force_read_decrement() };
505 assert_eq!(lock.reader_count(), 0);
506
507 unsafe { lock.force_read_decrement() };
508 assert_eq!(lock.reader_count(), 0);
509 assert!(lock.try_write().is_some());
510 }
511
512 #[test]
513 fn concurrent_readers_and_writers_preserve_updates() {
514 const THREADS: usize = 4;
515 const ITERS: usize = 2_000;
516
517 let lock = Arc::new(RwLock::new(0usize));
518 let observed = Arc::new(AtomicUsize::new(0));
519 let mut handles = Vec::new();
520
521 for _ in 0..THREADS {
522 let lock = lock.clone();
523 handles.push(thread::spawn(move || {
524 for _ in 0..ITERS {
525 *lock.write() += 1;
526 }
527 }));
528 }
529
530 for _ in 0..THREADS {
531 let lock = lock.clone();
532 let observed = observed.clone();
533 handles.push(thread::spawn(move || {
534 for _ in 0..ITERS {
535 let value = *lock.read();
536 observed.fetch_max(value, Ordering::Relaxed);
537 }
538 }));
539 }
540
541 for handle in handles {
542 handle.join().unwrap();
543 }
544
545 assert_eq!(*lock.read(), THREADS * ITERS);
546 assert!(observed.load(Ordering::Relaxed) <= THREADS * ITERS);
547 }
548}
549
550#[cfg(all(axtest, feature = "axtest"))]
551pub fn rwlock_constants_hold_for_test() -> bool {
552 assert!(READER == 1);
554 assert!(WRITER == 1 << (usize::BITS - 1));
555 assert!(MAX_READER == 1 << (usize::BITS - 2));
556
557 assert!(WRITER > READER);
559 assert!(MAX_READER == WRITER / 2);
561
562 true
563}
564
565#[cfg(all(axtest, feature = "axtest"))]
566pub fn rwlock_state_logic_hold_for_test() -> bool {
567 let idle: usize = 0;
571 assert!(idle & WRITER == 0); assert!(idle / READER == 0); let one_reader = READER;
576 assert!(one_reader & WRITER == 0); assert!(one_reader / READER == 1); let two_readers = 2 * READER;
581 assert!(two_readers & WRITER == 0); assert!(two_readers / READER == 2); let writer_only = WRITER;
586 assert!(writer_only & WRITER != 0); assert!(writer_only % READER == 0); let writer_one_reader = WRITER + READER;
591 assert!(writer_one_reader & WRITER != 0); let max_readers = MAX_READER * READER;
595 assert!(max_readers < WRITER); assert!(max_readers / READER == MAX_READER);
597
598 true
599}
600
601#[cfg(all(axtest, feature = "axtest"))]
602pub(crate) fn rwlock_constants_and_phantom_hold_for_test() -> bool {
603 assert_eq!(READER, 1);
605 assert!(WRITER > MAX_READER);
606 assert!(MAX_READER > 0);
607
608 use core::marker::PhantomData;
610 let _phantom: PhantomData<()> = PhantomData;
611
612 true
613}
614
615#[cfg(all(axtest, feature = "axtest"))]
616pub(crate) fn rwlock_state_transitions_hold_for_test() -> bool {
617 let unlocked: usize = 0;
621 assert!(unlocked == 0);
622
623 let one_reader = READER;
625 assert!(one_reader == 1);
626
627 let writer_only = WRITER;
629 assert!(writer_only != 0);
630
631 true
632}
633
634#[cfg(all(axtest, feature = "axtest"))]
635pub(crate) fn rwlock_guard_types_hold_for_test() -> bool {
636 true
640}
641
642#[cfg(all(axtest, feature = "axtest"))]
643pub(crate) fn rwlock_lockdep_and_feature_config_hold_for_test() -> bool {
644 #[cfg(feature = "lockdep")]
646 {
647 let _acquire = LockdepAcquire;
649 }
650
651 #[cfg(not(feature = "lockdep"))]
652 {
653 let acquire = LockdepAcquire;
655 acquire.finish(true);
656 acquire.finish(false);
657 }
658
659 assert_eq!(WRITER, 1usize << (usize::BITS - 1));
661
662 assert_eq!(MAX_READER, 1usize << (usize::BITS - 2));
664
665 true
666}
667
668#[cfg(all(axtest, feature = "axtest"))]
669pub(crate) fn rwlock_reader_writer_state_combinations_hold_for_test() -> bool {
670 let empty: usize = 0;
674 assert!(empty & WRITER == 0);
675 assert!(empty & !WRITER == 0); let one_r = READER;
679 assert!(one_r & WRITER == 0); assert!(one_r == 1);
681
682 let two_r = 2 * READER;
684 assert!(two_r & WRITER == 0);
685 assert!(two_r == 2);
686
687 let w_only = WRITER;
689 assert!(w_only & WRITER != 0); assert!(w_only & !WRITER == 0); let max_r = MAX_READER;
694 assert!(max_r & WRITER == 0);
695 assert!(max_r > 0);
696
697 true
698}