#![allow(unsafe_code)]
use parking_lot::Mutex as ParkingMutex;
use std::cell::UnsafeCell;
use std::future::Future;
use std::marker::PhantomData;
use std::mem::ManuallyDrop;
use std::ops::{Deref, DerefMut};
use std::pin::Pin;
use std::ptr::NonNull;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::{Context, Poll, Waker};
use crate::cx::Cx;
use crate::sync::lock_ordering::{self, LockRank};
use crate::time::Sleep;
use crate::types::Time;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum LockError {
Poisoned,
Cancelled,
TimedOut(Time),
PolledAfterCompletion,
}
impl std::fmt::Display for LockError {
#[inline]
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Poisoned => write!(f, "mutex poisoned"),
Self::Cancelled => write!(f, "mutex lock cancelled"),
Self::TimedOut(deadline) => write!(f, "mutex lock timed out at {deadline:?}"),
Self::PolledAfterCompletion => write!(f, "mutex future polled after completion"),
}
}
}
impl std::error::Error for LockError {}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum TryLockError {
Locked,
Poisoned,
}
impl std::fmt::Display for TryLockError {
#[inline]
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Locked => write!(f, "mutex is locked"),
Self::Poisoned => write!(f, "mutex poisoned"),
}
}
}
impl std::error::Error for TryLockError {}
#[derive(Debug)]
pub struct Mutex<T> {
data: UnsafeCell<T>,
poisoned: AtomicBool,
state: ParkingMutex<MutexState>,
name: &'static str,
rank: Option<LockRank>,
}
unsafe impl<T: Send> Send for Mutex<T> {}
unsafe impl<T: Send> Sync for Mutex<T> {}
#[derive(Debug)]
struct MutexState {
locked: bool,
waiters: WaiterChain,
granted_waiter: Option<WaiterId>,
}
use super::waiter::{WaiterChain, WaiterId};
impl<T> Mutex<T> {
#[inline]
#[must_use]
pub fn with_name(name: &'static str, value: T) -> Self {
let rank = lock_ordering::rank_for_lock_name(name);
Self {
data: UnsafeCell::new(value),
poisoned: AtomicBool::new(false),
state: ParkingMutex::new(MutexState {
locked: false,
waiters: WaiterChain::new(),
granted_waiter: None,
}),
name,
rank,
}
}
#[inline]
#[must_use]
pub fn new(value: T) -> Self {
Self::with_name("unknown", value)
}
#[inline]
#[must_use]
pub fn is_poisoned(&self) -> bool {
self.poisoned.load(Ordering::Acquire)
}
#[inline]
#[must_use]
pub fn is_locked(&self) -> bool {
self.state.lock().locked
}
#[inline]
#[must_use]
pub fn waiters(&self) -> usize {
self.state.lock().waiters.len()
}
#[inline]
pub fn lock<'a, 'b, Caps>(&'a self, cx: &'b Cx<Caps>) -> LockFuture<'a, 'b, T, Caps> {
LockFuture {
mutex: self,
cx,
waiter_id: None,
deadline_sleep: None,
completed: false,
}
}
#[inline]
pub fn lock_until<'a, 'b, Caps>(
&'a self,
cx: &'b Cx<Caps>,
deadline: Time,
) -> LockFuture<'a, 'b, T, Caps>
where
Caps: crate::cx::cap::HasTime,
{
LockFuture {
mutex: self,
cx,
waiter_id: None,
deadline_sleep: Some(cx.timer_driver().map_or_else(
|| Sleep::new(deadline),
|timer| Sleep::with_timer_driver(deadline, timer),
)),
completed: false,
}
}
#[inline]
pub fn try_lock(&self) -> Result<MutexGuard<'_, T>, TryLockError> {
let mut state = self.state.lock();
if self.is_poisoned() {
return Err(TryLockError::Poisoned);
}
if state.locked || state.granted_waiter.is_some() || !state.waiters.is_empty() {
return Err(TryLockError::Locked);
}
if let Some(rank) = self.rank {
lock_ordering::check_acquire(self.name, rank);
}
state.locked = true;
drop(state);
let lock_order = lock_ordering::record_guard_acquire(self.name, self.rank);
Ok(MutexGuard {
mutex: self,
lock_order,
_not_send: PhantomData,
})
}
#[inline]
pub fn try_lock_owned(self: &Arc<Self>) -> Result<OwnedMutexGuard<T>, TryLockError> {
OwnedMutexGuard::try_lock(Arc::clone(self))
}
#[inline]
pub fn get_mut(&mut self) -> Result<&mut T, LockError> {
if self.is_poisoned() {
return Err(LockError::Poisoned);
}
Ok(self.data.get_mut())
}
#[inline]
pub fn into_inner(self) -> Result<T, LockError> {
if self.is_poisoned() {
return Err(LockError::Poisoned);
}
Ok(self.data.into_inner())
}
#[inline]
fn poison(&self) {
self.poisoned.store(true, Ordering::Release);
}
#[cfg(any(test, feature = "test-internals"))]
#[doc(hidden)]
#[inline]
pub fn poison_for_testing(&self) {
self.poison();
}
#[inline]
fn unlock(&self, mut lock_order: lock_ordering::GuardLockOrder) {
let granted = {
let mut state = self.state.lock();
state.locked = false;
if let Some((id, waker, _)) = state.waiters.pop_front() {
state.granted_waiter = Some(id);
Some((id, waker))
} else {
state.granted_waiter = None;
None
}
};
lock_ordering::record_guard_release(&mut lock_order);
if let Some((id, waker)) = granted {
self.wake_granted(id, waker);
}
}
fn wake_granted(&self, id: crate::sync::waiter::WaiterId, waker: Waker) {
let mut pending = Some((id, waker));
let mut first_panic: Option<Box<dyn std::any::Any + Send>> = None;
while let Some((failed_id, waker)) = pending.take() {
match std::panic::catch_unwind(std::panic::AssertUnwindSafe(move || waker.wake())) {
Ok(()) => break,
Err(payload) => {
if first_panic.is_none() {
first_panic = Some(payload);
}
let mut state = self.state.lock();
if state.granted_waiter == Some(failed_id) {
state.granted_waiter = None;
if !state.locked {
if let Some((next_id, next_waker, _)) = state.waiters.pop_front() {
state.granted_waiter = Some(next_id);
pending = Some((next_id, next_waker));
}
}
}
}
}
}
if let Some(payload) = first_panic {
if !std::thread::panicking() {
std::panic::resume_unwind(payload);
}
}
}
}
impl<T: Default> Default for Mutex<T> {
#[inline]
fn default() -> Self {
Self::new(T::default())
}
}
pub struct LockFuture<'a, 'b, T, Caps = crate::cx::cap::All> {
mutex: &'a Mutex<T>,
cx: &'b Cx<Caps>,
waiter_id: Option<crate::sync::waiter::WaiterId>,
deadline_sleep: Option<Sleep>,
completed: bool,
}
impl<T, Caps> LockFuture<'_, '_, T, Caps> {
#[inline]
fn poll_deadline_sleep(&mut self, context: &mut Context<'_>) -> Option<Time> {
let sleep = self.deadline_sleep.as_mut()?;
let deadline = sleep.deadline();
match Pin::new(&mut *sleep).poll(context) {
Poll::Ready(()) => Some(deadline),
Poll::Pending => None,
}
}
#[inline]
fn grant_next_waiter(state: &mut MutexState) -> Option<(crate::sync::waiter::WaiterId, Waker)> {
if let Some((id, waker, _)) = state.waiters.pop_front() {
state.granted_waiter = Some(id);
Some((id, waker))
} else {
state.granted_waiter = None;
None
}
}
#[inline]
fn cleanup_waiter(&mut self) {
if let Some(waiter_id) = self.waiter_id.take() {
let (waker_to_wake, retired_waker) = {
let mut state = self.mutex.state.lock();
if state.granted_waiter == Some(waiter_id) {
state.granted_waiter = None;
if !state.locked {
(Self::grant_next_waiter(&mut state), None)
} else {
(None, None)
}
} else {
let is_head = state.waiters.front_id() == Some(waiter_id);
let retired_waker = state.waiters.remove(waiter_id);
let waker_to_wake =
if !state.locked && state.granted_waiter.is_none() && is_head {
Self::grant_next_waiter(&mut state)
} else {
None
};
(waker_to_wake, retired_waker)
}
};
if let Some((id, waker)) = waker_to_wake {
self.mutex.wake_granted(id, waker);
}
drop(retired_waker);
}
}
}
impl<'a, T, Caps> Future for LockFuture<'a, '_, T, Caps> {
type Output = Result<MutexGuard<'a, T>, LockError>;
#[inline]
#[allow(clippy::if_not_else, clippy::option_if_let_else)]
fn poll(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Self::Output> {
if self.completed {
return Poll::Ready(Err(LockError::PolledAfterCompletion));
}
if let Err(_e) = self.cx.checkpoint() {
self.completed = true;
self.cleanup_waiter();
return Poll::Ready(Err(LockError::Cancelled));
}
if let Some(deadline) = self.poll_deadline_sleep(context) {
self.completed = true;
self.cleanup_waiter();
return Poll::Ready(Err(LockError::TimedOut(deadline)));
}
let mut queued_waker = None;
loop {
let mut state = self.mutex.state.lock();
if self.mutex.is_poisoned() {
self.completed = true;
drop(state);
self.cleanup_waiter();
return Poll::Ready(Err(LockError::Poisoned));
}
if let Some(waiter_id) = self.waiter_id {
if state.granted_waiter == Some(waiter_id) {
if !state.locked {
if let Some(rank) = self.mutex.rank {
lock_ordering::check_acquire(self.mutex.name, rank);
}
state.granted_waiter = None;
state.locked = true;
self.waiter_id = None;
self.completed = true;
let lock_order =
lock_ordering::record_guard_acquire(self.mutex.name, self.mutex.rank);
return Poll::Ready(Ok(MutexGuard {
mutex: self.mutex,
lock_order,
_not_send: PhantomData,
}));
}
if queued_waker.is_none() {
drop(state);
queued_waker = Some(context.waker().clone());
continue;
}
state.granted_waiter = None;
let new_id = state.waiters.push_front_tagged(
queued_waker.take().expect("contended path cloned a waker"),
(),
);
drop(state);
self.waiter_id = Some(new_id);
return Poll::Pending;
}
}
if !state.locked && state.granted_waiter.is_none() && self.waiter_id.is_none() {
if let Some(rank) = self.mutex.rank {
lock_ordering::check_acquire(self.mutex.name, rank);
}
state.locked = true;
self.completed = true;
let lock_order =
lock_ordering::record_guard_acquire(self.mutex.name, self.mutex.rank);
return Poll::Ready(Ok(MutexGuard {
mutex: self.mutex,
lock_order,
_not_send: PhantomData,
}));
}
if queued_waker.is_none() {
drop(state);
queued_waker = Some(context.waker().clone());
continue;
}
let retired_waker = if let Some(waiter_id) = self.waiter_id {
let new_waker = queued_waker.take().expect("contended path cloned a waker");
match state.waiters.replace_waker(waiter_id, new_waker) {
Ok(retired_waker) => Some(retired_waker),
Err(new_waker) => {
let new_id = state.waiters.push_front_tagged(new_waker, ());
self.waiter_id = Some(new_id);
None
}
}
} else {
let id = state.waiters.push_back_tagged(
queued_waker.take().expect("contended path cloned a waker"),
(),
);
self.waiter_id = Some(id);
None
};
drop(state);
drop(retired_waker);
break;
}
if let Some(deadline) = self.poll_deadline_sleep(context) {
self.completed = true;
self.cleanup_waiter();
return Poll::Ready(Err(LockError::TimedOut(deadline)));
}
Poll::Pending
}
}
impl<T, Caps> Drop for LockFuture<'_, '_, T, Caps> {
fn drop(&mut self) {
self.cleanup_waiter();
}
}
#[must_use = "guard will be immediately released if not held"]
pub struct MutexGuard<'a, T> {
mutex: &'a Mutex<T>,
lock_order: lock_ordering::GuardLockOrder,
_not_send: PhantomData<*mut ()>,
}
unsafe impl<T: Sync> Sync for MutexGuard<'_, T> {}
impl<T: std::fmt::Debug> std::fmt::Debug for MutexGuard<'_, T> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MutexGuard").field("data", &**self).finish()
}
}
impl<T> Deref for MutexGuard<'_, T> {
type Target = T;
#[inline]
fn deref(&self) -> &T {
unsafe { &*self.mutex.data.get() }
}
}
impl<T> DerefMut for MutexGuard<'_, T> {
#[inline]
fn deref_mut(&mut self) -> &mut T {
unsafe { &mut *self.mutex.data.get() }
}
}
impl<T> Drop for MutexGuard<'_, T> {
fn drop(&mut self) {
if std::thread::panicking() {
self.mutex.poison();
}
self.mutex
.unlock(lock_ordering::take_guard_lock_order(&mut self.lock_order));
}
}
impl<'a, T> MutexGuard<'a, T> {
#[inline]
pub fn map<U: ?Sized, F>(mut self, f: F) -> MappedMutexGuard<'a, T, U>
where
F: FnOnce(&mut T) -> &mut U,
{
let data = NonNull::from(f(&mut *self));
let mutex = self.mutex;
let lock_order = lock_ordering::take_guard_lock_order(&mut self.lock_order);
let _guard = ManuallyDrop::new(self);
MappedMutexGuard {
mutex,
data,
lock_order,
_marker: PhantomData,
}
}
#[inline]
pub fn try_map<U: ?Sized, F>(mut self, f: F) -> Result<MappedMutexGuard<'a, T, U>, Self>
where
F: FnOnce(&mut T) -> Option<&mut U>,
{
let data = f(&mut *self).map(NonNull::from);
if let Some(data) = data {
let mutex = self.mutex;
let lock_order = lock_ordering::take_guard_lock_order(&mut self.lock_order);
let _guard = ManuallyDrop::new(self);
Ok(MappedMutexGuard {
mutex,
data,
lock_order,
_marker: PhantomData,
})
} else {
Err(self)
}
}
}
#[must_use = "guard will be immediately released if not held"]
pub struct MappedMutexGuard<'a, T, U: ?Sized> {
mutex: &'a Mutex<T>,
data: NonNull<U>,
lock_order: lock_ordering::GuardLockOrder,
_marker: PhantomData<&'a mut U>,
}
unsafe impl<T, U: ?Sized + Sync> Sync for MappedMutexGuard<'_, T, U> {}
impl<T, U: ?Sized + std::fmt::Debug> std::fmt::Debug for MappedMutexGuard<'_, T, U> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MappedMutexGuard")
.field("data", &&**self)
.finish()
}
}
impl<T, U: ?Sized> Deref for MappedMutexGuard<'_, T, U> {
type Target = U;
#[inline]
fn deref(&self) -> &U {
unsafe { self.data.as_ref() }
}
}
impl<T, U: ?Sized> DerefMut for MappedMutexGuard<'_, T, U> {
#[inline]
fn deref_mut(&mut self) -> &mut U {
unsafe { self.data.as_mut() }
}
}
impl<T, U: ?Sized> Drop for MappedMutexGuard<'_, T, U> {
fn drop(&mut self) {
if std::thread::panicking() {
self.mutex.poison();
}
self.mutex
.unlock(lock_ordering::take_guard_lock_order(&mut self.lock_order));
}
}
impl<'a, T, U: ?Sized> MappedMutexGuard<'a, T, U> {
#[inline]
pub fn map<V: ?Sized, F>(mut self, f: F) -> MappedMutexGuard<'a, T, V>
where
F: FnOnce(&mut U) -> &mut V,
{
let data = NonNull::from(f(&mut *self));
let mutex = self.mutex;
let lock_order = lock_ordering::take_guard_lock_order(&mut self.lock_order);
let _guard = ManuallyDrop::new(self);
MappedMutexGuard {
mutex,
data,
lock_order,
_marker: PhantomData,
}
}
#[inline]
pub fn try_map<V: ?Sized, F>(mut self, f: F) -> Result<MappedMutexGuard<'a, T, V>, Self>
where
F: FnOnce(&mut U) -> Option<&mut V>,
{
let data = f(&mut *self).map(NonNull::from);
if let Some(data) = data {
let mutex = self.mutex;
let lock_order = lock_ordering::take_guard_lock_order(&mut self.lock_order);
let _guard = ManuallyDrop::new(self);
Ok(MappedMutexGuard {
mutex,
data,
lock_order,
_marker: PhantomData,
})
} else {
Err(self)
}
}
}
#[must_use = "guard will be immediately released if not held"]
pub struct OwnedMutexGuard<T> {
mutex: Arc<Mutex<T>>,
lock_order: lock_ordering::GuardLockOrder,
}
unsafe impl<T: Send> Send for OwnedMutexGuard<T> {}
unsafe impl<T: Sync> Sync for OwnedMutexGuard<T> {}
impl<T> OwnedMutexGuard<T> {
pub async fn lock<Caps>(mutex: Arc<Mutex<T>>, cx: &Cx<Caps>) -> Result<Self, LockError> {
let mut borrowed_guard = mutex.as_ref().lock(cx).await?;
let lock_order = lock_ordering::take_guard_lock_order(&mut borrowed_guard.lock_order);
let _borrowed_guard = std::mem::ManuallyDrop::new(borrowed_guard);
Ok(Self { mutex, lock_order })
}
#[inline]
pub fn try_lock(mutex: Arc<Mutex<T>>) -> Result<Self, TryLockError> {
{
let mut state = mutex.state.lock();
if mutex.is_poisoned() {
return Err(TryLockError::Poisoned);
}
if state.locked || state.granted_waiter.is_some() || !state.waiters.is_empty() {
return Err(TryLockError::Locked);
}
if let Some(rank) = mutex.rank {
lock_ordering::check_acquire(mutex.name, rank);
}
state.locked = true;
}
let lock_order = lock_ordering::record_guard_acquire(mutex.name, mutex.rank);
Ok(Self { mutex, lock_order })
}
#[inline]
pub fn map<U: ?Sized, F>(mut self, f: F) -> OwnedMappedMutexGuard<T, U>
where
F: FnOnce(&mut T) -> &mut U,
{
let data = NonNull::from(f(&mut *self));
let mutex = unsafe { std::ptr::read(&self.mutex) };
let lock_order = lock_ordering::take_guard_lock_order(&mut self.lock_order);
let _guard = ManuallyDrop::new(self);
OwnedMappedMutexGuard {
mutex,
data,
lock_order,
_marker: PhantomData,
}
}
#[inline]
pub fn try_map<U: ?Sized, F>(mut self, f: F) -> Result<OwnedMappedMutexGuard<T, U>, Self>
where
F: FnOnce(&mut T) -> Option<&mut U>,
{
let data = f(&mut *self).map(NonNull::from);
if let Some(data) = data {
let mutex = unsafe { std::ptr::read(&self.mutex) };
let lock_order = lock_ordering::take_guard_lock_order(&mut self.lock_order);
let _guard = ManuallyDrop::new(self);
Ok(OwnedMappedMutexGuard {
mutex,
data,
lock_order,
_marker: PhantomData,
})
} else {
Err(self)
}
}
}
impl<T> Deref for OwnedMutexGuard<T> {
type Target = T;
#[inline]
fn deref(&self) -> &T {
unsafe { &*self.mutex.data.get() }
}
}
impl<T> DerefMut for OwnedMutexGuard<T> {
#[inline]
fn deref_mut(&mut self) -> &mut T {
unsafe { &mut *self.mutex.data.get() }
}
}
impl<T> Drop for OwnedMutexGuard<T> {
fn drop(&mut self) {
if std::thread::panicking() {
self.mutex.poison();
}
self.mutex
.unlock(lock_ordering::take_guard_lock_order(&mut self.lock_order));
}
}
#[must_use = "guard will be immediately released if not held"]
pub struct OwnedMappedMutexGuard<T, U: ?Sized> {
mutex: Arc<Mutex<T>>,
data: NonNull<U>,
lock_order: lock_ordering::GuardLockOrder,
_marker: PhantomData<*mut U>,
}
unsafe impl<T: Send, U: ?Sized + Send> Send for OwnedMappedMutexGuard<T, U> {}
unsafe impl<T: Send, U: ?Sized + Sync> Sync for OwnedMappedMutexGuard<T, U> {}
impl<T, U: ?Sized + std::fmt::Debug> std::fmt::Debug for OwnedMappedMutexGuard<T, U> {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("OwnedMappedMutexGuard")
.field("data", &&**self)
.finish()
}
}
impl<T, U: ?Sized> Deref for OwnedMappedMutexGuard<T, U> {
type Target = U;
#[inline]
fn deref(&self) -> &U {
unsafe { self.data.as_ref() }
}
}
impl<T, U: ?Sized> DerefMut for OwnedMappedMutexGuard<T, U> {
#[inline]
fn deref_mut(&mut self) -> &mut U {
unsafe { self.data.as_mut() }
}
}
impl<T, U: ?Sized> Drop for OwnedMappedMutexGuard<T, U> {
fn drop(&mut self) {
if std::thread::panicking() {
self.mutex.poison();
}
self.mutex
.unlock(lock_ordering::take_guard_lock_order(&mut self.lock_order));
}
}
impl<T, U: ?Sized> OwnedMappedMutexGuard<T, U> {
#[inline]
pub fn map<V: ?Sized, F>(mut self, f: F) -> OwnedMappedMutexGuard<T, V>
where
F: FnOnce(&mut U) -> &mut V,
{
let data = NonNull::from(f(&mut *self));
let mutex = unsafe { std::ptr::read(&self.mutex) };
let lock_order = lock_ordering::take_guard_lock_order(&mut self.lock_order);
let _guard = ManuallyDrop::new(self);
OwnedMappedMutexGuard {
mutex,
data,
lock_order,
_marker: PhantomData,
}
}
#[inline]
pub fn try_map<V: ?Sized, F>(mut self, f: F) -> Result<OwnedMappedMutexGuard<T, V>, Self>
where
F: FnOnce(&mut U) -> Option<&mut V>,
{
let data = f(&mut *self).map(NonNull::from);
if let Some(data) = data {
let mutex = unsafe { std::ptr::read(&self.mutex) };
let lock_order = lock_ordering::take_guard_lock_order(&mut self.lock_order);
let _guard = ManuallyDrop::new(self);
Ok(OwnedMappedMutexGuard {
mutex,
data,
lock_order,
_marker: PhantomData,
})
} else {
Err(self)
}
}
}
#[cfg(test)]
include!("mutex_tests.rs");