use core::cell::RefCell;
use core::cell::UnsafeCell;
use core::future::Future;
use core::marker::PhantomData;
use core::ops::{Deref, DerefMut};
use core::pin::Pin;
use core::task::Waker;
use core::task::{Context, Poll};
struct MutexState {
locked: bool,
waker: Option<Waker>,
}
pub struct Mutex<T> {
value: UnsafeCell<T>,
state: RefCell<MutexState>,
_no_send_sync: PhantomData<*mut T>, }
impl<T> Mutex<T> {
pub fn new(value: T) -> Self {
Mutex {
value: UnsafeCell::new(value),
state: RefCell::new(MutexState {
locked: false,
waker: None,
}),
_no_send_sync: PhantomData,
}
}
pub async fn lock(&self) -> MutexGuard<'_, T> {
LockFuture { mutex: self }.await;
MutexGuard { mutex: self }
}
pub fn try_lock(&self) -> Option<MutexGuard<'_, T>> {
let mut state = self.state.borrow_mut();
if !state.locked {
state.locked = true;
Some(MutexGuard { mutex: self })
} else {
None
}
}
pub fn get_mut(&mut self) -> &mut T {
self.value.get_mut()
}
pub unsafe fn read(&self) -> &T {
&*self.value.get()
}
}
pub struct MutexGuard<'a, T> {
mutex: &'a Mutex<T>,
}
impl<T> Deref for MutexGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
unsafe { &*self.mutex.value.get() }
}
}
impl<T> DerefMut for MutexGuard<'_, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
unsafe { &mut *self.mutex.value.get() }
}
}
impl<T> Drop for MutexGuard<'_, T> {
fn drop(&mut self) {
let mut mutex_state = self.mutex.state.borrow_mut();
mutex_state.locked = false;
if let Some(waker) = mutex_state.waker.take() {
waker.wake()
}
}
}
struct LockFuture<'a, T> {
mutex: &'a Mutex<T>,
}
impl<T> Future for LockFuture<'_, T> {
type Output = ();
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let mut mutex_state = self.mutex.state.borrow_mut();
if mutex_state.locked {
let new_waker = cx.waker();
match &mut mutex_state.waker {
Some(waker) if waker.will_wake(new_waker) => {
waker.clone_from(new_waker);
}
waker @ Some(_) => {
waker.take().unwrap().wake();
*waker = Some(new_waker.clone());
}
waker @ None => *waker = Some(new_waker.clone()),
};
Poll::Pending
} else {
mutex_state.locked = true;
Poll::Ready(())
}
}
}
#[cfg(test)]
mod tests {
use pollster::FutureExt as _;
use crate::sync::{join::join, select::select, yield_now};
use super::Mutex;
#[test]
pub fn test_mutex_no_concurrency() {
async {
let mut mutex = Mutex::new(0usize);
{
let mut guard = mutex.lock().await;
*guard += 1;
assert_eq!(*guard, 1, "The guard should be readable");
}
assert_eq!(
*mutex.get_mut(),
1,
"The internal mutex should have been updated"
)
}
.block_on()
}
#[test]
pub fn test_mutex_select_concurrency() {
async {
let mut mutex = Mutex::new(0usize);
for _ in 0..100 {
select(
async {
let mut guard = mutex.lock().await;
*guard += 1;
},
async {
let mut guard = mutex.lock().await;
*guard += 1;
},
)
.await;
}
assert_eq!(*mutex.get_mut(), 100);
}
.block_on()
}
#[test]
pub fn test_mutex_join_concurrency() {
async {
let mut mutex = Mutex::new(0usize);
for _ in 0..100 {
join(
async {
let mut guard = mutex.lock().await;
*guard += 1;
},
async {
let mut guard = mutex.lock().await;
*guard += 1;
},
)
.await;
}
assert_eq!(*mutex.get_mut(), 200);
}
.block_on()
}
#[test]
pub fn test_try_lock() {
async {
let mut mutex = Mutex::new(0usize);
join(
async {
let mut guard = mutex.lock().await;
for _ in 0..10 {
*guard += 1;
yield_now::yield_now().await;
}
},
async {
let mut i = 0;
loop {
if let Some(mut guard) = mutex.try_lock() {
*guard += 1;
break;
}
if i == 20 {
panic!("Try lock takes to long!");
}
i += 1;
yield_now::yield_now().await;
}
},
)
.await;
assert_eq!(*mutex.get_mut(), 11);
}
.block_on()
}
#[test]
pub fn test_drop_by_leaking() {
async {
let mutex = Mutex::new(Box::new(0));
let _guard = mutex.lock().await;
}
.block_on()
}
}