Skip to main content

ax_kspin/
rwlock.rs

1//! Spin-based read-write locks.
2
3use 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
38/// A spin-based read-write lock.
39///
40/// Readers may enter concurrently while a writer holds exclusive access. The
41/// lock never sleeps; failed acquisitions spin until the state changes. The
42/// guard `G` controls the atomic context used while the lock is held, matching
43/// [`BaseSpinLock`](crate::BaseSpinLock).
44pub 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
52/// A guard that provides shared data access.
53pub 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
62/// A guard that provides exclusive data access.
63pub 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    /// Creates a new [`BaseSpinRwLock`] wrapping the supplied data.
77    #[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    /// Consumes this lock and returns the underlying data.
90    #[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    /// Acquires a shared read lock, spinning until it is available.
168    #[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    /// Acquires an exclusive write lock, spinning until it is available.
190    #[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    /// Attempts to acquire a shared read lock.
212    #[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    /// Attempts to acquire an exclusive write lock.
236    #[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    /// Returns true if a writer currently holds the lock.
260    #[inline(always)]
261    pub fn is_write_locked(&self) -> bool {
262        self.state.load(Ordering::Acquire) & WRITER != 0
263    }
264
265    /// Returns the current reader count.
266    ///
267    /// This is only a heuristic; the value can change immediately after it is
268    /// loaded and must not be used for synchronization.
269    #[inline(always)]
270    pub fn reader_count(&self) -> usize {
271        // sync-lint: ignore suspicious_relaxed_mixed_ordering
272        self.state.load(Ordering::Relaxed) & !(WRITER | MAX_READER)
273    }
274
275    /// Returns the current writer count, which can only be 0 or 1.
276    ///
277    /// This is only a heuristic; the value can change immediately after it is
278    /// loaded and must not be used for synchronization.
279    #[inline(always)]
280    pub fn writer_count(&self) -> usize {
281        // sync-lint: ignore suspicious_relaxed_mixed_ordering
282        usize::from(self.state.load(Ordering::Relaxed) & WRITER != 0)
283    }
284
285    /// Force decrement the reader count.
286    ///
287    /// # Safety
288    ///
289    /// This is unsafe if called without a corresponding leaked read guard or if
290    /// any normal read guard is still expected to release that reader count.
291    /// If the reader count is already zero, this returns without changing the
292    /// state so a stale cleanup hook cannot underflow the lock and block future
293    /// writers permanently.
294    #[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    /// Force unlock exclusive write access.
323    ///
324    /// # Safety
325    ///
326    /// This is unsafe if called without a corresponding leaked write guard or
327    /// while readers are present.
328    #[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    /// Returns a mutable reference to the underlying data.
340    #[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    // RwLock state constants
553    assert!(READER == 1);
554    assert!(WRITER == 1 << (usize::BITS - 1));
555    assert!(MAX_READER == 1 << (usize::BITS - 2));
556
557    // WRITER should be much larger than READER
558    assert!(WRITER > READER);
559    // MAX_READER should be half of WRITER
560    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    // Test the state encoding logic
568
569    // No readers or writers: state = 0
570    let idle: usize = 0;
571    assert!(idle & WRITER == 0); // No writer bit set
572    assert!(idle / READER == 0); // Zero readers
573
574    // One reader: state = READER
575    let one_reader = READER;
576    assert!(one_reader & WRITER == 0); // No writer bit set
577    assert!(one_reader / READER == 1); // One reader
578
579    // Two readers: state = 2 * READER
580    let two_readers = 2 * READER;
581    assert!(two_readers & WRITER == 0); // No writer bit set
582    assert!(two_readers / READER == 2); // Two readers
583
584    // Writer present: state has WRITER bit set
585    let writer_only = WRITER;
586    assert!(writer_only & WRITER != 0); // Writer bit set
587    assert!(writer_only % READER == 0); // No reader count in lower bits
588
589    // Writer + one reader (theoretical)
590    let writer_one_reader = WRITER + READER;
591    assert!(writer_one_reader & WRITER != 0); // Writer bit set
592
593    // Max readers without overflow
594    let max_readers = MAX_READER * READER;
595    assert!(max_readers < WRITER); // Should not overlap with writer bit
596    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    // Test that constants are consistent
604    assert_eq!(READER, 1);
605    assert!(WRITER > MAX_READER);
606    assert!(MAX_READER > 0);
607
608    // Test PhantomData usage in BaseSpinRwLock
609    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    // Test state transitions for read-write lock
618
619    // Initial state (unlocked)
620    let unlocked: usize = 0;
621    assert!(unlocked == 0);
622
623    // One reader acquired
624    let one_reader = READER;
625    assert!(one_reader == 1);
626
627    // Writer acquired
628    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    // Test that guard types exist
637    // BaseSpinRwLockReadGuard and BaseSpinRwLockWriteGuard
638
639    true
640}
641
642#[cfg(all(axtest, feature = "axtest"))]
643pub(crate) fn rwlock_lockdep_and_feature_config_hold_for_test() -> bool {
644    // Test LockdepAcquire behavior based on feature flag
645    #[cfg(feature = "lockdep")]
646    {
647        // With lockdep feature, LockdepAcquire is crate::lockdep::Lockdep
648        let _acquire = LockdepAcquire;
649    }
650
651    #[cfg(not(feature = "lockdep"))]
652    {
653        // Without lockdep feature, LockdepAcquire is a simple unit struct
654        let acquire = LockdepAcquire;
655        acquire.finish(true);
656        acquire.finish(false);
657    }
658
659    // Test that WRITER bit position is correct
660    assert_eq!(WRITER, 1usize << (usize::BITS - 1));
661
662    // Test MAX_READER calculation
663    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    // Test various reader/writer state combinations
671
672    // No readers, no writer
673    let empty: usize = 0;
674    assert!(empty & WRITER == 0);
675    assert!(empty & !WRITER == 0); // No readers either
676
677    // One reader
678    let one_r = READER;
679    assert!(one_r & WRITER == 0); // No writer bit
680    assert!(one_r == 1);
681
682    // Two readers
683    let two_r = 2 * READER;
684    assert!(two_r & WRITER == 0);
685    assert!(two_r == 2);
686
687    // Writer only (no readers)
688    let w_only = WRITER;
689    assert!(w_only & WRITER != 0); // Writer bit set
690    assert!(w_only & !WRITER == 0); // No reader bits
691
692    // Max readers (without writer)
693    let max_r = MAX_READER;
694    assert!(max_r & WRITER == 0);
695    assert!(max_r > 0);
696
697    true
698}