Skip to main content

rustpython_common/lock/
thread_mutex.rs

1#![allow(clippy::needless_lifetimes)]
2
3use alloc::fmt;
4use core::{
5    cell::UnsafeCell,
6    marker::PhantomData,
7    ops::{Deref, DerefMut},
8    ptr::NonNull,
9    sync::atomic::{AtomicUsize, Ordering},
10};
11use lock_api::{GetThreadId, GuardNoSend, RawMutex};
12
13// based off ReentrantMutex from lock_api
14
15/// A mutex type that knows when it would deadlock
16pub struct RawThreadMutex<R: RawMutex, G: GetThreadId> {
17    owner: AtomicUsize,
18    mutex: R,
19    get_thread_id: G,
20}
21
22impl<R: RawMutex, G: GetThreadId> RawThreadMutex<R, G> {
23    #[allow(
24        clippy::declare_interior_mutable_const,
25        reason = "const initializer for lock primitive contains atomics by design"
26    )]
27    pub const INIT: Self = Self {
28        owner: AtomicUsize::new(0),
29        mutex: R::INIT,
30        get_thread_id: G::INIT,
31    };
32
33    #[inline]
34    fn lock_internal<F: FnOnce() -> bool>(&self, try_lock: F) -> Option<bool> {
35        let id = self.get_thread_id.nonzero_thread_id().get();
36        if self.owner.load(Ordering::Relaxed) == id {
37            return None;
38        }
39        if !try_lock() {
40            return Some(false);
41        }
42        self.owner.store(id, Ordering::Relaxed);
43        Some(true)
44    }
45
46    /// Blocks for the mutex to be available, and returns true if the mutex isn't already
47    /// locked on the current thread.
48    pub fn lock(&self) -> bool {
49        self.lock_internal(|| {
50            self.mutex.lock();
51            true
52        })
53        .is_some()
54    }
55
56    /// Like `lock()` but wraps the blocking wait in `wrap_fn`.
57    /// The caller can use this to detach thread state while waiting.
58    pub fn lock_wrapped<F: FnOnce(&dyn Fn())>(&self, wrap_fn: F) -> bool {
59        let id = self.get_thread_id.nonzero_thread_id().get();
60        if self.owner.load(Ordering::Relaxed) == id {
61            return false;
62        }
63        wrap_fn(&|| self.mutex.lock());
64        self.owner.store(id, Ordering::Relaxed);
65        true
66    }
67
68    /// Returns `Some(true)` if able to successfully lock without blocking, `Some(false)`
69    /// otherwise, and `None` when the mutex is already locked on the current thread.
70    pub fn try_lock(&self) -> Option<bool> {
71        self.lock_internal(|| self.mutex.try_lock())
72    }
73
74    /// Unlocks this mutex. The inner mutex may not be unlocked if
75    /// this mutex was acquired previously in the current thread.
76    ///
77    /// # Safety
78    ///
79    /// This method may only be called if the mutex is held by the current thread.
80    pub unsafe fn unlock(&self) {
81        self.owner.store(0, Ordering::Relaxed);
82        unsafe { self.mutex.unlock() };
83    }
84}
85
86impl<R: RawMutex, G: GetThreadId> RawThreadMutex<R, G> {
87    /// Reset this mutex to its initial (unlocked, unowned) state after `fork()`.
88    ///
89    /// # Safety
90    ///
91    /// Must only be called from the single-threaded child process immediately
92    /// after `fork()`, before any other thread is created.
93    #[cfg(unix)]
94    pub unsafe fn reinit_after_fork(&self) {
95        self.owner.store(0, Ordering::Relaxed);
96        unsafe {
97            let mutex_ptr = &self.mutex as *const R as *mut u8;
98            core::ptr::write_bytes(mutex_ptr, 0, core::mem::size_of::<R>());
99        }
100    }
101}
102
103unsafe impl<R: RawMutex + Send, G: GetThreadId + Send> Send for RawThreadMutex<R, G> {}
104unsafe impl<R: RawMutex + Sync, G: GetThreadId + Sync> Sync for RawThreadMutex<R, G> {}
105
106pub struct ThreadMutex<R: RawMutex, G: GetThreadId, T: ?Sized> {
107    raw: RawThreadMutex<R, G>,
108    data: UnsafeCell<T>,
109}
110
111impl<R: RawMutex, G: GetThreadId, T> ThreadMutex<R, G, T> {
112    pub const fn new(val: T) -> Self {
113        Self {
114            raw: RawThreadMutex::INIT,
115            data: UnsafeCell::new(val),
116        }
117    }
118
119    pub fn into_inner(self) -> T {
120        self.data.into_inner()
121    }
122}
123impl<R: RawMutex, G: GetThreadId, T: Default> Default for ThreadMutex<R, G, T> {
124    fn default() -> Self {
125        Self::new(T::default())
126    }
127}
128impl<R: RawMutex, G: GetThreadId, T> From<T> for ThreadMutex<R, G, T> {
129    fn from(val: T) -> Self {
130        Self::new(val)
131    }
132}
133impl<R: RawMutex, G: GetThreadId, T: ?Sized> ThreadMutex<R, G, T> {
134    /// Access the underlying raw thread mutex.
135    pub fn raw(&self) -> &RawThreadMutex<R, G> {
136        &self.raw
137    }
138
139    pub fn lock(&self) -> Option<ThreadMutexGuard<'_, R, G, T>> {
140        if self.raw.lock() {
141            Some(ThreadMutexGuard {
142                mu: self,
143                marker: PhantomData,
144            })
145        } else {
146            None
147        }
148    }
149
150    /// Like `lock()` but wraps the blocking wait in `wrap_fn`.
151    /// The caller can use this to detach thread state while waiting.
152    pub fn lock_wrapped<F: FnOnce(&dyn Fn())>(
153        &self,
154        wrap_fn: F,
155    ) -> Option<ThreadMutexGuard<'_, R, G, T>> {
156        if self.raw.lock_wrapped(wrap_fn) {
157            Some(ThreadMutexGuard {
158                mu: self,
159                marker: PhantomData,
160            })
161        } else {
162            None
163        }
164    }
165
166    pub fn try_lock(&self) -> Result<ThreadMutexGuard<'_, R, G, T>, TryLockThreadError> {
167        match self.raw.try_lock() {
168            Some(true) => Ok(ThreadMutexGuard {
169                mu: self,
170                marker: PhantomData,
171            }),
172            Some(false) => Err(TryLockThreadError::Other),
173            None => Err(TryLockThreadError::Current),
174        }
175    }
176}
177
178#[derive(Clone, Copy)]
179pub enum TryLockThreadError {
180    /// Failed to lock because mutex was already locked on another thread.
181    Other,
182    /// Failed to lock because mutex was already locked on current thread.
183    Current,
184}
185
186struct LockedPlaceholder(&'static str);
187
188impl fmt::Debug for LockedPlaceholder {
189    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
190        f.write_str(self.0)
191    }
192}
193
194impl<R: RawMutex, G: GetThreadId, T: ?Sized + fmt::Debug> fmt::Debug for ThreadMutex<R, G, T> {
195    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
196        match self.try_lock() {
197            Ok(guard) => f
198                .debug_struct("ThreadMutex")
199                .field("data", &&*guard)
200                .finish(),
201            Err(e) => {
202                let msg = match e {
203                    TryLockThreadError::Other => "<locked on other thread>",
204                    TryLockThreadError::Current => "<locked on current thread>",
205                };
206                f.debug_struct("ThreadMutex")
207                    .field("data", &LockedPlaceholder(msg))
208                    .finish()
209            }
210        }
211    }
212}
213
214unsafe impl<R: RawMutex + Send, G: GetThreadId + Send, T: ?Sized + Send> Send
215    for ThreadMutex<R, G, T>
216{
217}
218unsafe impl<R: RawMutex + Sync, G: GetThreadId + Sync, T: ?Sized + Send> Sync
219    for ThreadMutex<R, G, T>
220{
221}
222
223pub struct ThreadMutexGuard<'a, R: RawMutex, G: GetThreadId, T: ?Sized> {
224    mu: &'a ThreadMutex<R, G, T>,
225    marker: PhantomData<(&'a mut T, GuardNoSend)>,
226}
227impl<'a, R: RawMutex, G: GetThreadId, T: ?Sized> ThreadMutexGuard<'a, R, G, T> {
228    pub fn map<U, F: FnOnce(&mut T) -> &mut U>(
229        mut s: Self,
230        f: F,
231    ) -> MappedThreadMutexGuard<'a, R, G, U> {
232        let data = f(&mut s).into();
233        let mu = &s.mu.raw;
234        core::mem::forget(s);
235        MappedThreadMutexGuard {
236            mu,
237            data,
238            marker: PhantomData,
239        }
240    }
241    pub fn try_map<U, F: FnOnce(&mut T) -> Option<&mut U>>(
242        mut s: Self,
243        f: F,
244    ) -> Result<MappedThreadMutexGuard<'a, R, G, U>, Self> {
245        if let Some(data) = f(&mut s) {
246            let data = data.into();
247            let mu = &s.mu.raw;
248            core::mem::forget(s);
249            Ok(MappedThreadMutexGuard {
250                mu,
251                data,
252                marker: PhantomData,
253            })
254        } else {
255            Err(s)
256        }
257    }
258}
259impl<R: RawMutex, G: GetThreadId, T: ?Sized> Deref for ThreadMutexGuard<'_, R, G, T> {
260    type Target = T;
261    fn deref(&self) -> &T {
262        unsafe { &*self.mu.data.get() }
263    }
264}
265impl<R: RawMutex, G: GetThreadId, T: ?Sized> DerefMut for ThreadMutexGuard<'_, R, G, T> {
266    fn deref_mut(&mut self) -> &mut T {
267        unsafe { &mut *self.mu.data.get() }
268    }
269}
270impl<R: RawMutex, G: GetThreadId, T: ?Sized> Drop for ThreadMutexGuard<'_, R, G, T> {
271    fn drop(&mut self) {
272        unsafe { self.mu.raw.unlock() }
273    }
274}
275impl<R: RawMutex, G: GetThreadId, T: ?Sized + fmt::Display> fmt::Display
276    for ThreadMutexGuard<'_, R, G, T>
277{
278    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
279        fmt::Display::fmt(&**self, f)
280    }
281}
282impl<R: RawMutex, G: GetThreadId, T: ?Sized + fmt::Debug> fmt::Debug
283    for ThreadMutexGuard<'_, R, G, T>
284{
285    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
286        fmt::Debug::fmt(&**self, f)
287    }
288}
289pub struct MappedThreadMutexGuard<'a, R: RawMutex, G: GetThreadId, T: ?Sized> {
290    mu: &'a RawThreadMutex<R, G>,
291    data: NonNull<T>,
292    marker: PhantomData<(&'a mut T, GuardNoSend)>,
293}
294impl<'a, R: RawMutex, G: GetThreadId, T: ?Sized> MappedThreadMutexGuard<'a, R, G, T> {
295    pub fn map<U, F: FnOnce(&mut T) -> &mut U>(
296        mut s: Self,
297        f: F,
298    ) -> MappedThreadMutexGuard<'a, R, G, U> {
299        let data = f(&mut s).into();
300        let mu = s.mu;
301        core::mem::forget(s);
302        MappedThreadMutexGuard {
303            mu,
304            data,
305            marker: PhantomData,
306        }
307    }
308    pub fn try_map<U, F: FnOnce(&mut T) -> Option<&mut U>>(
309        mut s: Self,
310        f: F,
311    ) -> Result<MappedThreadMutexGuard<'a, R, G, U>, Self> {
312        if let Some(data) = f(&mut s) {
313            let data = data.into();
314            let mu = s.mu;
315            core::mem::forget(s);
316            Ok(MappedThreadMutexGuard {
317                mu,
318                data,
319                marker: PhantomData,
320            })
321        } else {
322            Err(s)
323        }
324    }
325}
326impl<R: RawMutex, G: GetThreadId, T: ?Sized> Deref for MappedThreadMutexGuard<'_, R, G, T> {
327    type Target = T;
328    fn deref(&self) -> &T {
329        unsafe { self.data.as_ref() }
330    }
331}
332impl<R: RawMutex, G: GetThreadId, T: ?Sized> DerefMut for MappedThreadMutexGuard<'_, R, G, T> {
333    fn deref_mut(&mut self) -> &mut T {
334        unsafe { self.data.as_mut() }
335    }
336}
337impl<R: RawMutex, G: GetThreadId, T: ?Sized> Drop for MappedThreadMutexGuard<'_, R, G, T> {
338    fn drop(&mut self) {
339        unsafe { self.mu.unlock() }
340    }
341}
342impl<R: RawMutex, G: GetThreadId, T: ?Sized + fmt::Display> fmt::Display
343    for MappedThreadMutexGuard<'_, R, G, T>
344{
345    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
346        fmt::Display::fmt(&**self, f)
347    }
348}
349impl<R: RawMutex, G: GetThreadId, T: ?Sized + fmt::Debug> fmt::Debug
350    for MappedThreadMutexGuard<'_, R, G, T>
351{
352    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
353        fmt::Debug::fmt(&**self, f)
354    }
355}