shuttle_std/sync/
mutex.rs1use crate::sync::{LockResult, PoisonError, TryLockError, TryLockResult};
2use crate::sync::{ResourceSignature, ResourceType};
3use shuttle_engine::current;
4use shuttle_engine::future::batch_semaphore::{BatchSemaphore, Fairness};
5use shuttle_engine::runtime::task::TaskId;
6use shuttle_engine::runtime::thread;
7use std::cell::RefCell;
8use std::fmt::{Debug, Display};
9use std::ops::{Deref, DerefMut};
10use std::panic::{RefUnwindSafe, UnwindSafe};
11use tracing::trace;
12
13pub struct Mutex<T: ?Sized> {
15 state: RefCell<MutexState>,
16 semaphore: BatchSemaphore,
17 inner: std::sync::Mutex<T>,
18}
19
20pub struct MutexGuard<'a, T: ?Sized> {
22 inner: Option<std::sync::MutexGuard<'a, T>>,
23 mutex: &'a Mutex<T>,
24}
25
26#[derive(Debug)]
27struct MutexState {
28 holder: Option<TaskId>,
29}
30
31impl<T> Mutex<T> {
32 #[track_caller]
34 pub const fn new(value: T) -> Self {
35 Self::new_internal(value, ResourceSignature::new_const(ResourceType::Mutex))
36 }
37
38 pub(crate) const fn new_internal(value: T, signature: ResourceSignature) -> Self {
39 let state = MutexState { holder: None };
40 Self {
41 state: RefCell::new(state),
42 semaphore: BatchSemaphore::const_new_with_signature(1, Fairness::Unfair, signature),
43 inner: std::sync::Mutex::new(value),
44 }
45 }
46}
47
48impl<T: ?Sized> Mutex<T> {
49 pub fn lock(&self) -> LockResult<MutexGuard<'_, T>> {
51 let me = current::me();
52
53 let mut state = self.state.borrow_mut();
54 trace!(holder=?state.holder, semaphore=?self.semaphore, "waiting to acquire mutex {:p}", self);
55 drop(state);
56
57 if !self.semaphore.is_closed() {
58 state = self.state.borrow_mut();
60 assert!(
61 match &state.holder {
62 Some(holder) => *holder != me,
63 None => true,
64 },
65 "deadlock! task {me:?} tried to acquire a Mutex it already holds"
66 );
67 drop(state);
68
69 self.semaphore.acquire_blocking(1).unwrap();
70 } else {
71 thread::switch();
73 }
74
75 state = self.state.borrow_mut();
76 assert!(state.holder.is_none());
77 state.holder = Some(me);
78 drop(state);
79
80 trace!(semaphore=?self.semaphore, "acquired mutex {:p}", self);
81
82 let result = match self.inner.try_lock() {
84 Ok(guard) => Ok(MutexGuard {
85 inner: Some(guard),
86 mutex: self,
87 }),
88 Err(TryLockError::Poisoned(guard)) => Err(PoisonError::new(MutexGuard {
89 inner: Some(guard.into_inner()),
90 mutex: self,
91 })),
92 Err(TryLockError::WouldBlock) => unreachable!("mutex state out of sync"),
93 };
94
95 result
96 }
97
98 pub fn try_lock(&self) -> TryLockResult<MutexGuard<'_, T>> {
103 let me = current::me();
104
105 let mut state = self.state.borrow_mut();
106 trace!(holder=?state.holder, semaphore=?self.semaphore, "trying to acquire mutex {:p}", self);
107 drop(state);
108
109 self.semaphore.try_acquire(1).map_err(|_| TryLockError::WouldBlock)?;
113
114 state = self.state.borrow_mut();
115 state.holder = Some(me);
116 drop(state);
117
118 trace!(semaphore=?self.semaphore, "acquired mutex {:p}", self);
119
120 let result = match self.inner.try_lock() {
122 Ok(guard) => Ok(MutexGuard {
123 inner: Some(guard),
124 mutex: self,
125 }),
126 Err(TryLockError::Poisoned(guard)) => Err(TryLockError::Poisoned(PoisonError::new(MutexGuard {
127 inner: Some(guard.into_inner()),
128 mutex: self,
129 }))),
130 Err(TryLockError::WouldBlock) => unreachable!("mutex state out of sync"),
131 };
132
133 result
134 }
135
136 #[inline]
141 pub fn get_mut(&mut self) -> LockResult<&mut T> {
142 self.inner.get_mut()
143 }
144
145 pub fn into_inner(self) -> LockResult<T>
147 where
148 T: Sized,
149 {
150 let state = self.state.borrow();
151 assert!(state.holder.is_none());
152
153 self.semaphore.try_acquire(1).unwrap();
155
156 self.inner.into_inner()
157 }
158
159 pub fn clear_poison(&self) {
161 self.inner.clear_poison()
162 }
163}
164
165unsafe impl<T: Send + ?Sized> Send for Mutex<T> {}
170unsafe impl<T: Send + ?Sized> Sync for Mutex<T> {}
171
172impl<T: ?Sized> UnwindSafe for Mutex<T> {}
174impl<T: ?Sized> RefUnwindSafe for Mutex<T> {}
175
176impl<T: Default> Default for Mutex<T> {
177 fn default() -> Self {
178 Self::new(Default::default())
179 }
180}
181
182impl<T: ?Sized + Debug> Debug for Mutex<T> {
183 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
184 Debug::fmt(&self.inner, f)
185 }
186}
187
188impl<'a, T: ?Sized> MutexGuard<'a, T> {
189 pub(super) fn unlock(self) -> &'a Mutex<T> {
191 self.mutex
192 }
193}
194
195impl<T: ?Sized> Drop for MutexGuard<'_, T> {
196 fn drop(&mut self) {
197 self.mutex.semaphore.release(1);
199
200 self.inner = None;
202
203 let mut state = self.mutex.state.borrow_mut();
204 trace!(semaphore=?self.mutex.semaphore, "releasing mutex {:p}", self.mutex);
205 state.holder = None;
206 }
207}
208
209impl<T: ?Sized> Deref for MutexGuard<'_, T> {
210 type Target = T;
211
212 fn deref(&self) -> &Self::Target {
213 self.inner.as_ref().unwrap()
214 }
215}
216
217impl<T: ?Sized> DerefMut for MutexGuard<'_, T> {
218 fn deref_mut(&mut self) -> &mut Self::Target {
219 self.inner.as_mut().unwrap()
220 }
221}
222
223impl<T: Debug + ?Sized> Debug for MutexGuard<'_, T> {
224 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
225 Debug::fmt(&self.inner.as_ref().unwrap(), f)
226 }
227}
228
229impl<T: Display + ?Sized> Display for MutexGuard<'_, T> {
230 fn fmt(&self, f: &mut std::fmt::Formatter) -> std::fmt::Result {
231 (**self).fmt(f)
232 }
233}
234
235impl<T> shuttle_engine::annotations::WithName for &Mutex<T> {
236 fn with_name_and_kind(self, name: Option<&str>, kind: Option<&str>) -> Self {
237 (&self.semaphore).with_name_and_kind(name, kind.or(Some("shuttle::sync::Mutex")));
238 self
239 }
240}
241
242impl<T> shuttle_engine::annotations::WithName for Mutex<T> {
243 fn with_name_and_kind(self, name: Option<&str>, kind: Option<&str>) -> Self {
244 (&self).with_name_and_kind(name, kind);
245 self
246 }
247}
248
249#[cfg(test)]
250mod tests {
251 use super::*;
252
253 #[test]
254 fn unique_resource_signature_mutex() {
255 shuttle_schedulers::check_random(
256 || {
257 let mutex1 = Mutex::new(0);
258 let mutex2 = Mutex::new(0);
259 assert_ne!(mutex1.semaphore.signature(), mutex2.semaphore.signature());
260 },
261 1,
262 );
263 }
264}