Skip to main content

shuttle_std/sync/
rwlock.rs

1use crate::sync::{ResourceSignature, ResourceType};
2use shuttle_engine::future::batch_semaphore::{BatchSemaphore, Fairness};
3use shuttle_engine::runtime::execution::ExecutionState;
4use shuttle_engine::runtime::task::{TaskId, TaskSet};
5use shuttle_engine::runtime::thread;
6use std::cell::RefCell;
7use std::fmt::{Debug, Display};
8use std::ops::{Deref, DerefMut};
9use std::panic::{RefUnwindSafe, UnwindSafe};
10use std::sync::{LockResult, PoisonError, TryLockError, TryLockResult};
11use tracing::trace;
12
13/// (Theoretical) max number of readers holding the same `RwLock`. Based on
14/// the `tokio` implementation.
15const MAX_READS: usize = (u32::MAX >> 3) as usize;
16
17/// A reader-writer lock, the same as [`std::sync::RwLock`].
18///
19/// Unlike [`std::sync::RwLock`], the same thread is never allowed to acquire the read side of a
20/// `RwLock` more than once. The `std` version is ambiguous about what behavior is allowed here, so
21/// we choose the most conservative one.
22pub struct RwLock<T: ?Sized> {
23    state: RefCell<RwLockState>,
24    semaphore: BatchSemaphore,
25    inner: std::sync::RwLock<T>,
26}
27
28#[derive(Debug)]
29struct RwLockState {
30    holder: RwLockHolder,
31}
32
33#[derive(PartialEq, Eq, Debug)]
34enum RwLockHolder {
35    Read(TaskSet),
36    Write(TaskId),
37    None,
38}
39
40#[derive(PartialEq, Eq, Debug, Clone, Copy)]
41enum RwLockType {
42    Read,
43    Write,
44}
45
46impl RwLockType {
47    /// Number of semaphore permits corresponding to the given lock type.
48    fn num_permits(&self) -> usize {
49        match self {
50            Self::Read => 1,
51            Self::Write => MAX_READS,
52        }
53    }
54}
55
56impl<T> RwLock<T> {
57    /// Create a new instance of an `RwLock<T>` which is unlocked.
58    #[track_caller]
59    pub const fn new(value: T) -> Self {
60        let state = RwLockState {
61            holder: RwLockHolder::None,
62        };
63
64        Self {
65            inner: std::sync::RwLock::new(value),
66            semaphore: BatchSemaphore::const_new_with_signature(
67                MAX_READS,
68                Fairness::Unfair,
69                ResourceSignature::new_const(ResourceType::RwLock),
70            ),
71            state: RefCell::new(state),
72        }
73    }
74}
75
76impl<T: ?Sized> RwLock<T> {
77    /// Locks this rwlock with shared read access, blocking the current thread until it can be
78    /// acquired.
79    pub fn read(&self) -> LockResult<RwLockReadGuard<'_, T>> {
80        self.lock(RwLockType::Read);
81
82        match self.inner.try_read() {
83            Ok(guard) => Ok(RwLockReadGuard {
84                inner: Some(guard),
85                rwlock: self,
86                me: ExecutionState::me(),
87            }),
88            Err(TryLockError::Poisoned(err)) => Err(PoisonError::new(RwLockReadGuard {
89                inner: Some(err.into_inner()),
90                rwlock: self,
91                me: ExecutionState::me(),
92            })),
93            Err(TryLockError::WouldBlock) => panic!("rwlock state out of sync"),
94        }
95    }
96
97    /// Locks this rwlock with exclusive write access, blocking the current thread until it can
98    /// be acquired.
99    pub fn write(&self) -> LockResult<RwLockWriteGuard<'_, T>> {
100        self.lock(RwLockType::Write);
101
102        match self.inner.try_write() {
103            Ok(guard) => Ok(RwLockWriteGuard {
104                inner: Some(guard),
105                rwlock: self,
106                me: ExecutionState::me(),
107            }),
108            Err(TryLockError::Poisoned(err)) => Err(PoisonError::new(RwLockWriteGuard {
109                inner: Some(err.into_inner()),
110                rwlock: self,
111                me: ExecutionState::me(),
112            })),
113            Err(TryLockError::WouldBlock) => panic!("rwlock state out of sync"),
114        }
115    }
116
117    /// Attempts to acquire this rwlock with shared read access.
118    ///
119    /// If the access could not be granted at this time, then Err is returned. This function does
120    /// not block.
121    ///
122    /// Note that unlike [`std::sync::RwLock::try_read`], if the current thread already holds this
123    /// read lock, `try_read` will return Err.
124    pub fn try_read(&self) -> TryLockResult<RwLockReadGuard<'_, T>> {
125        if self.try_lock(RwLockType::Read) {
126            match self.inner.try_read() {
127                Ok(guard) => Ok(RwLockReadGuard {
128                    inner: Some(guard),
129                    rwlock: self,
130                    me: ExecutionState::me(),
131                }),
132                Err(TryLockError::Poisoned(err)) => Err(TryLockError::Poisoned(PoisonError::new(RwLockReadGuard {
133                    inner: Some(err.into_inner()),
134                    rwlock: self,
135                    me: ExecutionState::me(),
136                }))),
137                Err(TryLockError::WouldBlock) => panic!("rwlock state out of sync"),
138            }
139        } else {
140            Err(TryLockError::WouldBlock)
141        }
142    }
143
144    /// Attempts to acquire this rwlock with shared read access.
145    ///
146    /// If the access could not be granted at this time, then Err is returned. This function does
147    /// not block.
148    pub fn try_write(&self) -> TryLockResult<RwLockWriteGuard<'_, T>> {
149        if self.try_lock(RwLockType::Write) {
150            match self.inner.try_write() {
151                Ok(guard) => Ok(RwLockWriteGuard {
152                    inner: Some(guard),
153                    rwlock: self,
154                    me: ExecutionState::me(),
155                }),
156                Err(TryLockError::Poisoned(err)) => Err(TryLockError::Poisoned(PoisonError::new(RwLockWriteGuard {
157                    inner: Some(err.into_inner()),
158                    rwlock: self,
159                    me: ExecutionState::me(),
160                }))),
161                Err(TryLockError::WouldBlock) => panic!("rwlock state out of sync"),
162            }
163        } else {
164            Err(TryLockError::WouldBlock)
165        }
166    }
167
168    /// Returns a mutable reference to the underlying data.
169    ///
170    /// Since this call borrows the `RwLock` mutably, no actual locking needs to
171    /// take place---the mutable borrow statically guarantees no locks exist.
172    #[inline]
173    pub fn get_mut(&mut self) -> LockResult<&mut T> {
174        self.inner.get_mut()
175    }
176
177    /// Consumes this `RwLock`, returning the underlying data
178    pub fn into_inner(self) -> LockResult<T>
179    where
180        T: Sized,
181    {
182        let state = self.state.borrow();
183        assert_eq!(state.holder, RwLockHolder::None);
184
185        // Update the receiver's clock with the RwLock clock
186        self.semaphore.try_acquire(MAX_READS).unwrap();
187
188        self.inner.into_inner()
189    }
190
191    /// Acquire the lock in the provided mode, blocking this thread until it succeeds.
192    fn lock(&self, typ: RwLockType) {
193        let me = ExecutionState::me();
194
195        let mut state = self.state.borrow_mut();
196        trace!(
197            holder = ?state.holder,
198            semaphore = ?self.semaphore,
199            "acquiring {:?} lock on rwlock {:p}",
200            typ,
201            self,
202        );
203        drop(state);
204
205        if !self.semaphore.is_closed() {
206            // Detect deadlock due to re-entrancy.
207            state = self.state.borrow_mut();
208            assert!(
209                match &state.holder {
210                    RwLockHolder::Write(writer) => *writer != me,
211                    RwLockHolder::Read(readers) => !readers.contains(me),
212                    RwLockHolder::None => true,
213                },
214                "deadlock! task {me:?} tried to acquire a RwLock it already holds"
215            );
216            drop(state);
217
218            self.semaphore.acquire_blocking(typ.num_permits()).unwrap();
219        } else {
220            // we always need to allow for a context switch to make the previous event visible for completeness
221            thread::switch();
222        }
223
224        state = self.state.borrow_mut();
225        match (typ, &mut state.holder) {
226            (RwLockType::Write, RwLockHolder::None) => {
227                state.holder = RwLockHolder::Write(me);
228            }
229            (RwLockType::Read, RwLockHolder::None) => {
230                let mut readers = TaskSet::new();
231                readers.insert(me);
232                state.holder = RwLockHolder::Read(readers);
233            }
234            (RwLockType::Read, RwLockHolder::Read(readers)) => {
235                assert!(readers.insert(me));
236            }
237            _ => {
238                panic!(
239                    "resumed a waiting {:?} thread while the lock was in state {:?}",
240                    typ, state.holder
241                );
242            }
243        }
244        trace!(
245            holder = ?state.holder,
246            semaphore = ?self.semaphore,
247            "acquired {:?} lock on rwlock {:p}",
248            typ,
249            self
250        );
251        drop(state);
252    }
253
254    /// Attempt to acquire this lock in the provided mode, but without blocking. Returns `true` if
255    /// the lock was able to be acquired without blocking, or `false` otherwise.
256    fn try_lock(&self, typ: RwLockType) -> bool {
257        let me = ExecutionState::me();
258
259        let mut state = self.state.borrow_mut();
260        trace!(
261            holder = ?state.holder,
262            semaphore = ?self.semaphore,
263            "trying to acquire {:?} lock on rwlock {:p}",
264            typ,
265            self,
266        );
267        drop(state);
268
269        // Semaphore is never closed, so an error here is always `NoPermits`.
270        let mut acquired = self.semaphore.try_acquire(typ.num_permits()).is_ok();
271        if acquired {
272            state = self.state.borrow_mut();
273            match (typ, &mut state.holder) {
274                (RwLockType::Write, RwLockHolder::None) => {
275                    state.holder = RwLockHolder::Write(me);
276                }
277                (RwLockType::Read, RwLockHolder::None) => {
278                    let mut readers = TaskSet::new();
279                    readers.insert(me);
280                    state.holder = RwLockHolder::Read(readers);
281                }
282                (RwLockType::Read, RwLockHolder::Read(readers)) => {
283                    // If we already hold the read lock, `insert` returns false, which will cause this
284                    // acquisition to fail with `WouldBlock` so we can diagnose potential deadlocks.
285                    acquired = readers.insert(me);
286                }
287                _ => (),
288            };
289            drop(state);
290        }
291
292        trace!(
293            "{} {:?} lock on rwlock {:p}",
294            if acquired { "acquired" } else { "failed to acquire" },
295            typ,
296            self,
297        );
298
299        acquired
300    }
301
302    /// Clear the poisoned state from a lock.
303    #[inline]
304    pub fn clear_poison(&self) {
305        self.inner.clear_poison();
306    }
307}
308
309// Safety: RwLock is never actually passed across true threads, only across continuations. The
310// Rc<RefCell<_>> type therefore can't be preempted mid-bookkeeping-operation.
311// TODO we shouldn't need to do this, but RefCell is not Send, and anything we put within a RwLock
312// TODO needs to be Send.
313unsafe impl<T: Send + ?Sized> Send for RwLock<T> {}
314unsafe impl<T: Send + ?Sized> Sync for RwLock<T> {}
315
316// TODO this is the RefCell biting us again
317impl<T: ?Sized> UnwindSafe for RwLock<T> {}
318impl<T: ?Sized> RefUnwindSafe for RwLock<T> {}
319
320impl<T: Default> Default for RwLock<T> {
321    fn default() -> Self {
322        Self::new(Default::default())
323    }
324}
325
326impl<T: ?Sized + Debug> Debug for RwLock<T> {
327    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
328        Debug::fmt(&self.inner, f)
329    }
330}
331
332/// RAII structure used to release the shared read access of a `RwLock` when dropped.
333pub struct RwLockReadGuard<'a, T: ?Sized> {
334    inner: Option<std::sync::RwLockReadGuard<'a, T>>,
335    rwlock: &'a RwLock<T>,
336    me: TaskId,
337}
338
339impl<T: ?Sized> Deref for RwLockReadGuard<'_, T> {
340    type Target = T;
341
342    fn deref(&self) -> &Self::Target {
343        self.inner.as_ref().unwrap().deref()
344    }
345}
346
347impl<T: Debug> Debug for RwLockReadGuard<'_, T> {
348    fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
349        Debug::fmt(&self.inner.as_ref().unwrap(), f)
350    }
351}
352
353impl<T: Display + ?Sized> Display for RwLockReadGuard<'_, T> {
354    fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
355        (**self).fmt(f)
356    }
357}
358
359impl<T: ?Sized> Drop for RwLockReadGuard<'_, T> {
360    fn drop(&mut self) {
361        // Execution teardown can unwind the task from the yield point in `release` (see
362        // `ExecutionState::tear_down`). `release` releases the permit even then, and the rest of
363        // unlocking has to happen too, so that the destructors that teardown runs later can lock the
364        // `RwLock`.
365        struct Unlock<'g, 'a, T: ?Sized>(&'g mut RwLockReadGuard<'a, T>);
366        impl<T: ?Sized> Drop for Unlock<'_, '_, T> {
367            #[inline]
368            fn drop(&mut self) {
369                let guard = &mut *self.0;
370                guard.inner = None;
371
372                let mut state = guard.rwlock.state.borrow_mut();
373                trace!(
374                    holder = ?state.holder,
375                    semaphore = ?guard.rwlock.semaphore,
376                    "releasing Read lock on rwlock {:p}",
377                    guard.rwlock
378                );
379                let RwLockHolder::Read(readers) = &mut state.holder else {
380                    panic!("exiting a reader but rwlock is in the wrong state {:?}", state.holder);
381                };
382                assert!(readers.remove(guard.me));
383                if readers.is_empty() {
384                    state.holder = RwLockHolder::None;
385                }
386                drop(state);
387            }
388        }
389        let unlock = Unlock(self);
390
391        unlock.0.rwlock.semaphore.release(RwLockType::Read.num_permits());
392    }
393}
394
395/// RAII structure used to release the exclusive write access of a `RwLock` when dropped.
396pub struct RwLockWriteGuard<'a, T: ?Sized> {
397    inner: Option<std::sync::RwLockWriteGuard<'a, T>>,
398    rwlock: &'a RwLock<T>,
399    me: TaskId,
400}
401
402impl<T: ?Sized> Deref for RwLockWriteGuard<'_, T> {
403    type Target = T;
404
405    fn deref(&self) -> &Self::Target {
406        self.inner.as_ref().unwrap().deref()
407    }
408}
409
410impl<T: ?Sized> DerefMut for RwLockWriteGuard<'_, T> {
411    fn deref_mut(&mut self) -> &mut Self::Target {
412        self.inner.as_mut().unwrap().deref_mut()
413    }
414}
415
416impl<T: Debug> Debug for RwLockWriteGuard<'_, T> {
417    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
418        Debug::fmt(&self.inner.as_ref().unwrap(), f)
419    }
420}
421
422impl<T: Display + ?Sized> Display for RwLockWriteGuard<'_, T> {
423    fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
424        (**self).fmt(f)
425    }
426}
427
428impl<T: ?Sized> Drop for RwLockWriteGuard<'_, T> {
429    fn drop(&mut self) {
430        // As for `RwLockReadGuard`.
431        struct Unlock<'g, 'a, T: ?Sized>(&'g mut RwLockWriteGuard<'a, T>);
432        impl<T: ?Sized> Drop for Unlock<'_, '_, T> {
433            #[inline]
434            fn drop(&mut self) {
435                let guard = &mut *self.0;
436                // Teardown unwinding the task is no panic, which mustn't poison the inner lock.
437                let clear_poison = std::thread::panicking()
438                    && !guard.rwlock.inner.is_poisoned()
439                    && ExecutionState::unwinding_for_teardown();
440                guard.inner = None;
441                if clear_poison {
442                    guard.rwlock.inner.clear_poison();
443                }
444
445                let mut state = guard.rwlock.state.borrow_mut();
446                trace!(
447                    holder = ?state.holder,
448                    semaphore = ?guard.rwlock.semaphore,
449                    "releasing Write lock on rwlock {:p}",
450                    guard.rwlock
451                );
452                assert_eq!(state.holder, RwLockHolder::Write(guard.me));
453                state.holder = RwLockHolder::None;
454                drop(state);
455            }
456        }
457        let unlock = Unlock(self);
458
459        unlock.0.rwlock.semaphore.release(RwLockType::Write.num_permits());
460    }
461}
462
463#[cfg(test)]
464mod tests {
465    use super::*;
466
467    #[test]
468    fn unique_resource_signature_rwlock() {
469        shuttle_schedulers::check_random(
470            || {
471                let rwlock1 = RwLock::new(0);
472                let rwlock2 = RwLock::new(0);
473                assert_ne!(rwlock1.semaphore.signature(), rwlock2.semaphore.signature());
474            },
475            1,
476        );
477    }
478}