moirai_async/sync/
mutex.rs1#![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
14pub struct Mutex<T> {
16 data: UnsafeCell<T>,
17 state: std::sync::Mutex<MutexState>,
18}
19
20unsafe impl<T: Send + Sync> Sync for Mutex<T> {}
25unsafe impl<T: Send> Send for Mutex<T> {}
28
29struct MutexState {
30 locked: bool,
31 waiters: WaitQueue<()>,
32}
33
34impl<T> Mutex<T> {
35 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 pub fn lock(&self) -> MutexLockFuture<'_, T> {
48 MutexLockFuture {
49 mutex: self,
50 id: None,
51 }
52 }
53
54 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 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
96pub 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
147pub 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 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 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}