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
13pub 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 pub fn lock(&self) -> bool {
49 self.lock_internal(|| {
50 self.mutex.lock();
51 true
52 })
53 .is_some()
54 }
55
56 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 pub fn try_lock(&self) -> Option<bool> {
71 self.lock_internal(|| self.mutex.try_lock())
72 }
73
74 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 #[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 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 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 Other,
182 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}