#[cfg(feature = "std")]
pub(crate) use std::sync::MutexGuard;
#[cfg(not(feature = "std"))]
pub(crate) use self::spin::MutexGuard;
#[derive(Default)]
pub(crate) struct Mutex<T> {
#[cfg(feature = "std")]
inner: std::sync::Mutex<T>,
#[cfg(not(feature = "std"))]
inner: spin::Mutex<T>,
}
impl<T> Mutex<T> {
pub(crate) fn lock(&self) -> MutexGuard<'_, T> {
#[cfg(feature = "std")]
{
self.inner.lock().unwrap_or_else(|err| err.into_inner())
}
#[cfg(not(feature = "std"))]
{
self.inner.lock()
}
}
}
#[cfg(not(feature = "std"))]
mod spin {
use core::cell::UnsafeCell;
use core::ops::{Deref, DerefMut};
use core::sync::atomic::{AtomicBool, Ordering};
#[derive(Default)]
pub(crate) struct Mutex<T> {
locked: AtomicBool,
value: UnsafeCell<T>,
}
unsafe impl<T: Send> Sync for Mutex<T> {}
impl<T> Mutex<T> {
pub(crate) fn lock(&self) -> MutexGuard<'_, T> {
while self
.locked
.compare_exchange_weak(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_err()
{
core::hint::spin_loop();
}
MutexGuard(self)
}
}
pub(crate) struct MutexGuard<'a, T>(&'a Mutex<T>);
impl<T> Deref for MutexGuard<'_, T> {
type Target = T;
fn deref(&self) -> &T {
unsafe { &*self.0.value.get() }
}
}
impl<T> DerefMut for MutexGuard<'_, T> {
fn deref_mut(&mut self) -> &mut T {
unsafe { &mut *self.0.value.get() }
}
}
impl<T> Drop for MutexGuard<'_, T> {
fn drop(&mut self) {
self.0.locked.store(false, Ordering::Release);
}
}
}
#[test]
fn test_mutex() {
use alloc::sync::Arc;
use alloc::vec::Vec;
let mutex = Arc::new(Mutex::<Vec<usize>>::default());
let threads = (0..4)
.map(|idx| {
let mutex = mutex.clone();
std::thread::spawn(move || {
for _ in 0..if cfg!(miri) { 10 } else { 1000 } {
mutex.lock().push(idx);
}
})
})
.collect::<Vec<_>>();
for thread in threads {
thread.join().unwrap();
}
let values = mutex.lock();
assert_eq!(values.len(), 4 * if cfg!(miri) { 10 } else { 1000 });
}