use super::wait_queue::WaitQueue;
use std::cell::UnsafeCell;
use std::ops::{Deref, DerefMut};
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll};
#[repr(align(64))]
pub struct Mutex<T: ?Sized> {
locked: AtomicBool,
wait: WaitQueue,
data: UnsafeCell<T>,
}
unsafe impl<T: ?Sized + Send> Send for Mutex<T> {}
unsafe impl<T: ?Sized + Send> Sync for Mutex<T> {}
impl<T> Mutex<T> {
#[must_use]
pub const fn new(data: T) -> Self {
Self {
locked: AtomicBool::new(false),
wait: WaitQueue::new(),
data: UnsafeCell::new(data),
}
}
#[inline(always)]
pub fn into_inner(self) -> T {
self.data.into_inner()
}
}
impl<T: ?Sized> Mutex<T> {
#[inline(always)]
pub async fn lock(&self) -> MutexGuard<'_, T> {
std::future::poll_fn(|cx| self.poll_lock(cx)).await
}
#[must_use]
#[inline(always)]
pub fn try_lock(&self) -> Option<MutexGuard<'_, T>> {
self.locked
.compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
.then(|| MutexGuard { mutex: self })
}
#[inline]
fn poll_lock(&self, cx: &Context<'_>) -> Poll<MutexGuard<'_, T>> {
if !self.wait.has_waiters()
&& let Some(guard) = self.try_lock()
{
return Poll::Ready(guard);
}
let token = self.wait.register(cx.waker());
if let Some(guard) = self.try_lock() {
self.wait.cancel(token);
return Poll::Ready(guard);
}
Poll::Pending
}
#[inline(always)]
pub const fn get_mut(&mut self) -> &mut T {
self.data.get_mut()
}
}
impl<T: Default> Default for Mutex<T> {
fn default() -> Self {
Self {
locked: AtomicBool::new(false),
wait: WaitQueue::new(),
data: UnsafeCell::new(T::default()),
}
}
}
#[repr(align(64))]
pub struct MutexGuard<'a, T: ?Sized> {
mutex: &'a Mutex<T>,
}
unsafe impl<T: ?Sized + Send> Send for MutexGuard<'_, T> {}
unsafe impl<T: ?Sized + Sync> Sync for MutexGuard<'_, T> {}
impl<T: ?Sized> Deref for MutexGuard<'_, T> {
type Target = T;
#[inline(always)]
fn deref(&self) -> &T {
unsafe { &*self.mutex.data.get() }
}
}
impl<T: ?Sized> DerefMut for MutexGuard<'_, T> {
#[inline(always)]
fn deref_mut(&mut self) -> &mut T {
unsafe { &mut *self.mutex.data.get() }
}
}
impl<T: ?Sized> Drop for MutexGuard<'_, T> {
#[inline(always)]
fn drop(&mut self) {
self.mutex.locked.store(false, Ordering::Release);
self.mutex.wait.wake_one();
}
}