Skip to main content

moirai_async/sync/
mutex.rs

1#![expect(
2    clippy::unwrap_used,
3    reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
4)]
5
6use std::cell::UnsafeCell;
7use std::future::Future;
8use std::ops::{Deref, DerefMut};
9use std::pin::Pin;
10use std::task::{Context, Poll};
11
12use crate::sync::wait_queue::{WaitQueue, WaiterPoll};
13
14/// Async mutual-exclusion lock over `T`.
15pub struct Mutex<T> {
16    data: UnsafeCell<T>,
17    state: std::sync::Mutex<MutexState>,
18}
19
20// SAFETY: shared access serializes on `state`; a guard exists only after
21// acquiring it, so `&T`/`&mut T` from `data` are never concurrent. `Sync`
22// additionally needs `T: Sync` because guard derefs expose `&T` across
23// threads.
24unsafe impl<T: Send + Sync> Sync for Mutex<T> {}
25// SAFETY: the mutex owns its data and moves with it; no thread-local or
26// address-sensitive state exists beyond `T` itself.
27unsafe impl<T: Send> Send for Mutex<T> {}
28
29struct MutexState {
30    locked: bool,
31    waiters: WaitQueue<()>,
32}
33
34impl<T> Mutex<T> {
35    /// Create an unlocked mutex owning `data`.
36    pub fn new(data: T) -> Self {
37        Self {
38            data: UnsafeCell::new(data),
39            state: std::sync::Mutex::new(MutexState {
40                locked: false,
41                waiters: WaitQueue::new(),
42            }),
43        }
44    }
45
46    /// Acquire the lock, waiting for the current holder to release.
47    pub fn lock(&self) -> MutexLockFuture<'_, T> {
48        MutexLockFuture {
49            mutex: self,
50            id: None,
51        }
52    }
53
54    /// Acquire without waiting; `None` when already held.
55    pub fn try_lock(&self) -> Option<MutexGuard<'_, T>> {
56        let mut state = self.state.lock().unwrap();
57        if !state.locked {
58            state.locked = true;
59            Some(MutexGuard { mutex: self })
60        } else {
61            None
62        }
63    }
64
65    fn release(&self) {
66        // The waker leaves the state lock before it is woken: `Waker::wake` may
67        // poll the task inline on this thread, and that poll re-locks this
68        // state — waking under the lock would self-deadlock. Same discipline as
69        // `rwlock`'s release paths and `hybrid::notify`.
70        let waker = {
71            let mut state = self.state.lock().unwrap();
72            let waker = state.waiters.grant_oldest(());
73            if waker.is_none() {
74                state.locked = false;
75            }
76            waker
77        };
78        if let Some(waker) = waker {
79            waker.wake();
80        }
81    }
82}
83
84impl<T: Default> Default for Mutex<T> {
85    fn default() -> Self {
86        Self::new(T::default())
87    }
88}
89
90impl<T> From<T> for Mutex<T> {
91    fn from(data: T) -> Self {
92        Self::new(data)
93    }
94}
95
96/// Future returned by [`Mutex::lock`].
97pub struct MutexLockFuture<'a, T> {
98    mutex: &'a Mutex<T>,
99    id: Option<u64>,
100}
101
102impl<'a, T> Future for MutexLockFuture<'a, T> {
103    type Output = MutexGuard<'a, T>;
104
105    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
106        let mut state = self.mutex.state.lock().unwrap();
107
108        if let Some(id) = self.id {
109            match state.waiters.poll_waiter(id, cx.waker()) {
110                WaiterPoll::Granted(()) => {
111                    self.id = None;
112                    return Poll::Ready(MutexGuard { mutex: self.mutex });
113                }
114                WaiterPoll::Pending => return Poll::Pending,
115                WaiterPoll::NotRegistered => {}
116            }
117        }
118
119        if !state.locked {
120            state.locked = true;
121            if let Some(id) = self.id.take() {
122                let _removed = state.waiters.deregister(id);
123            }
124            return Poll::Ready(MutexGuard { mutex: self.mutex });
125        }
126
127        if self.id.is_none() {
128            self.id = Some(state.waiters.register(cx.waker().clone()));
129        }
130
131        Poll::Pending
132    }
133}
134
135impl<'a, T> Drop for MutexLockFuture<'a, T> {
136    fn drop(&mut self) {
137        if let Some(id) = self.id
138            && let Ok(mut state) = self.mutex.state.lock()
139            && state.waiters.deregister(id).is_some()
140        {
141            drop(state);
142            self.mutex.release();
143        }
144    }
145}
146
147/// Exclusive access guard; releases the lock on drop.
148pub struct MutexGuard<'a, T> {
149    pub(crate) mutex: &'a Mutex<T>,
150}
151
152impl<'a, T> Deref for MutexGuard<'a, T> {
153    type Target = T;
154    fn deref(&self) -> &Self::Target {
155        // SAFETY: the guard's existence proves the state lock was acquired;
156        // no other guard can coexist, so the shared reborrow is exclusive in
157        // practice and no mutable alias is live.
158        unsafe { &*self.mutex.data.get() }
159    }
160}
161
162impl<'a, T> DerefMut for MutexGuard<'a, T> {
163    fn deref_mut(&mut self) -> &mut Self::Target {
164        // SAFETY: unique guard plus serialized acquisition prove no other
165        // reference to `data` exists while this guard lives.
166        unsafe { &mut *self.mutex.data.get() }
167    }
168}
169
170impl<'a, T> Drop for MutexGuard<'a, T> {
171    fn drop(&mut self) {
172        self.mutex.release();
173    }
174}
175
176#[cfg(test)]
177mod tests {
178    use super::Mutex;
179    use std::future::Future;
180    use std::pin::Pin;
181    use std::task::{Context, Poll, Waker};
182
183    fn poll_future<F: Future + Unpin>(future: &mut F) -> Poll<F::Output> {
184        let mut context = Context::from_waker(Waker::noop());
185        Pin::new(future).poll(&mut context)
186    }
187
188    #[test]
189    fn test_mutex_lock_unlock() {
190        let lock = Mutex::new(42_u32);
191        let mut guard = lock.try_lock().expect("lock must succeed");
192        assert_eq!(*guard, 42);
193        *guard = 7;
194        drop(guard);
195        let guard = lock.try_lock().expect("lock must succeed after drop");
196        assert_eq!(*guard, 7);
197    }
198
199    #[test]
200    fn test_mutex_async_lock_release_grants_waiter() {
201        let lock = Mutex::new(10_u32);
202        let guard = lock.try_lock().expect("lock must succeed");
203        let mut waiter = lock.lock();
204        assert!(matches!(poll_future(&mut waiter), Poll::Pending));
205        drop(guard);
206        match poll_future(&mut waiter) {
207            Poll::Ready(mut guard) => *guard += 5,
208            Poll::Pending => panic!("waiter must be granted after release"),
209        }
210        let guard = lock.try_lock().expect("lock must succeed after waiter");
211        assert_eq!(*guard, 15);
212    }
213
214    #[test]
215    fn test_mutex_cancellation_safety() {
216        let lock = Mutex::new(0_u32);
217        let guard = lock.try_lock().expect("lock must succeed");
218        let mut waiter = lock.lock();
219        assert!(matches!(poll_future(&mut waiter), Poll::Pending));
220        drop(waiter);
221        drop(guard);
222        let guard = lock
223            .try_lock()
224            .expect("lock must be available after cancel+release");
225        assert_eq!(*guard, 0);
226    }
227
228    #[test]
229    fn test_mutex_cancellation_restores_permit() {
230        let lock = Mutex::new(0_u32);
231        let guard = lock.try_lock().expect("lock must succeed");
232        let mut waiter = lock.lock();
233        assert!(matches!(poll_future(&mut waiter), Poll::Pending));
234        drop(guard);
235        drop(waiter);
236        let guard = lock.try_lock().expect("lock must be available");
237        assert_eq!(*guard, 0);
238    }
239
240    #[test]
241    fn test_mutex_exclusive_access() {
242        let lock = Mutex::new(Vec::<i32>::new());
243        let guard = lock.try_lock().expect("lock must succeed");
244        let mut waiter = lock.lock();
245        assert!(matches!(poll_future(&mut waiter), Poll::Pending));
246        drop(guard);
247        match poll_future(&mut waiter) {
248            Poll::Ready(mut guard) => guard.push(1),
249            Poll::Pending => panic!("waiter must be granted"),
250        }
251        assert!(lock.try_lock().is_some());
252    }
253}