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        self.rwlock.semaphore.release(RwLockType::Read.num_permits());
362
363        self.inner = None;
364
365        let mut state = self.rwlock.state.borrow_mut();
366        trace!(
367            holder = ?state.holder,
368            semaphore = ?self.rwlock.semaphore,
369            "releasing Read lock on rwlock {:p}",
370            self.rwlock
371        );
372        let RwLockHolder::Read(readers) = &mut state.holder else {
373            panic!("exiting a reader but rwlock is in the wrong state {:?}", state.holder);
374        };
375        assert!(readers.remove(self.me));
376        if readers.is_empty() {
377            state.holder = RwLockHolder::None;
378        }
379        drop(state);
380    }
381}
382
383/// RAII structure used to release the exclusive write access of a `RwLock` when dropped.
384pub struct RwLockWriteGuard<'a, T: ?Sized> {
385    inner: Option<std::sync::RwLockWriteGuard<'a, T>>,
386    rwlock: &'a RwLock<T>,
387    me: TaskId,
388}
389
390impl<T: ?Sized> Deref for RwLockWriteGuard<'_, T> {
391    type Target = T;
392
393    fn deref(&self) -> &Self::Target {
394        self.inner.as_ref().unwrap().deref()
395    }
396}
397
398impl<T: ?Sized> DerefMut for RwLockWriteGuard<'_, T> {
399    fn deref_mut(&mut self) -> &mut Self::Target {
400        self.inner.as_mut().unwrap().deref_mut()
401    }
402}
403
404impl<T: Debug> Debug for RwLockWriteGuard<'_, T> {
405    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
406        Debug::fmt(&self.inner.as_ref().unwrap(), f)
407    }
408}
409
410impl<T: Display + ?Sized> Display for RwLockWriteGuard<'_, T> {
411    fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
412        (**self).fmt(f)
413    }
414}
415
416impl<T: ?Sized> Drop for RwLockWriteGuard<'_, T> {
417    fn drop(&mut self) {
418        self.rwlock.semaphore.release(RwLockType::Write.num_permits());
419
420        self.inner = None;
421
422        let mut state = self.rwlock.state.borrow_mut();
423        trace!(
424            holder = ?state.holder,
425            semaphore = ?self.rwlock.semaphore,
426            "releasing Write lock on rwlock {:p}",
427            self.rwlock
428        );
429        assert_eq!(state.holder, RwLockHolder::Write(self.me));
430        state.holder = RwLockHolder::None;
431        drop(state);
432    }
433}
434
435#[cfg(test)]
436mod tests {
437    use super::*;
438
439    #[test]
440    fn unique_resource_signature_rwlock() {
441        shuttle_schedulers::check_random(
442            || {
443                let rwlock1 = RwLock::new(0);
444                let rwlock2 = RwLock::new(0);
445                assert_ne!(rwlock1.semaphore.signature(), rwlock2.semaphore.signature());
446            },
447            1,
448        );
449    }
450}