use crate::core::detector;
use crate::core::locks::{
NEXT_LOCK_ID,
contention::{ContentionState, SlowWaiter},
};
use crate::core::types::{LockId, ThreadId, get_current_thread_id};
#[cfg(feature = "logging-and-visualization")]
use crate::core::{Events, logger};
use parking_lot::{Mutex as ParkingLotMutex, MutexGuard as ParkingLotMutexGuard};
use std::ops::{Deref, DerefMut};
use std::sync::atomic::{AtomicUsize, Ordering};
pub struct Mutex<T> {
id: LockId,
inner: ParkingLotMutex<T>,
creator_thread_id: ThreadId,
state: MutexState,
}
struct MutexState {
owner: AtomicUsize,
contention: ContentionState,
}
pub struct MutexGuard<'a, T> {
thread_id: ThreadId,
lock_id: LockId,
guard: ParkingLotMutexGuard<'a, T>,
state: &'a MutexState,
tracked_globally: bool,
}
impl<T> Mutex<T> {
pub fn new(value: T) -> Self {
let id = NEXT_LOCK_ID.fetch_add(1, Ordering::SeqCst);
let creator_thread_id = get_current_thread_id();
detector::mutex::create_mutex(id, Some(creator_thread_id));
Mutex {
id,
inner: ParkingLotMutex::new(value),
creator_thread_id,
state: MutexState {
owner: AtomicUsize::new(0),
contention: ContentionState::new(),
},
}
}
pub fn id(&self) -> LockId {
self.id
}
pub fn creator_thread_id(&self) -> ThreadId {
self.creator_thread_id
}
pub fn lock(&self) -> MutexGuard<'_, T> {
let thread_id = get_current_thread_id();
let tid_usize = thread_id;
#[cfg(not(feature = "stress-test"))]
if let Some(guard) = self.inner.try_lock() {
self.state.owner.store(tid_usize, Ordering::Release);
let tracked_globally =
cfg!(feature = "lock-order-graph") || self.state.contention.has_waiters();
#[cfg(feature = "logging-and-visualization")]
{
if logger::LOGGING_ENABLED.load(Ordering::Relaxed) {
logger::log_interaction_event(thread_id, self.id, Events::MutexAttempt);
}
}
if tracked_globally {
detector::mutex::complete_acquire(thread_id, self.id);
}
#[cfg(feature = "logging-and-visualization")]
{
if logger::LOGGING_ENABLED.load(Ordering::Relaxed) {
logger::log_interaction_event(thread_id, self.id, Events::MutexAcquired);
}
}
return MutexGuard {
thread_id,
lock_id: self.id,
guard,
state: &self.state,
tracked_globally,
};
}
let slow_waiter = self.state.contention.register();
let (rechecked_guard, deadlock_info) = detector::mutex::acquire_slow_with_recheck(
thread_id,
self.id,
|| self.inner.try_lock(),
|| {
let owner = self.state.owner.load(Ordering::Acquire);
(owner != 0).then_some(owner as ThreadId)
},
);
if let Some(info) = deadlock_info {
detector::deadlock_handling::process_deadlock(info);
}
if let Some(guard) = rechecked_guard {
self.state.owner.store(tid_usize, Ordering::Release);
drop(slow_waiter);
return MutexGuard {
thread_id,
lock_id: self.id,
guard,
state: &self.state,
tracked_globally: true,
};
}
let guard = self.inner.lock();
self.state.owner.store(tid_usize, Ordering::Release);
detector::mutex::complete_acquire(thread_id, self.id);
drop(slow_waiter);
MutexGuard {
thread_id,
lock_id: self.id,
guard,
state: &self.state,
tracked_globally: true,
}
}
pub fn try_lock(&self) -> Option<MutexGuard<'_, T>> {
let thread_id = get_current_thread_id();
let tid_usize = thread_id;
if let Some(guard) = self.inner.try_lock() {
self.state.owner.store(tid_usize, Ordering::Release);
let tracked_globally =
cfg!(feature = "lock-order-graph") || self.state.contention.has_waiters();
#[cfg(feature = "logging-and-visualization")]
{
if logger::LOGGING_ENABLED.load(Ordering::Relaxed) {
logger::log_interaction_event(thread_id, self.id, Events::MutexAttempt);
}
}
if tracked_globally {
detector::mutex::complete_acquire(thread_id, self.id);
}
#[cfg(feature = "logging-and-visualization")]
{
if logger::LOGGING_ENABLED.load(Ordering::Relaxed) {
logger::log_interaction_event(thread_id, self.id, Events::MutexAcquired);
}
}
Some(MutexGuard {
thread_id,
lock_id: self.id,
guard,
state: &self.state,
tracked_globally,
})
} else {
None
}
}
pub fn into_inner(self) -> T
where
T: Sized,
{
detector::mutex::destroy_mutex(self.id);
let mutex = std::mem::ManuallyDrop::new(self);
unsafe { std::ptr::read(&mutex.inner) }.into_inner()
}
pub fn get_mut(&mut self) -> &mut T {
self.inner.get_mut()
}
}
impl<T> Drop for Mutex<T> {
fn drop(&mut self) {
detector::mutex::destroy_mutex(self.id);
}
}
impl<T> Deref for MutexGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
self.guard.deref()
}
}
impl<T> DerefMut for MutexGuard<'_, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
self.guard.deref_mut()
}
}
impl<'a, T> MutexGuard<'a, T> {
pub(crate) fn inner_guard(&mut self) -> &mut ParkingLotMutexGuard<'a, T> {
&mut self.guard
}
pub(crate) fn lock_id(&self) -> LockId {
self.lock_id
}
pub(crate) fn register_condvar_waiter(&self) -> SlowWaiter<'a> {
self.state.contention.register()
}
pub(crate) fn clear_ownership(&self) {
self.state.owner.store(0, Ordering::Release);
}
pub(crate) fn restore_ownership(&self) {
self.state.owner.store(self.thread_id, Ordering::Release);
}
pub(crate) fn mark_tracked_globally(&mut self) {
self.tracked_globally = true;
}
#[cfg(all(test, not(feature = "lock-order-graph")))]
pub(crate) fn is_tracked_globally(&self) -> bool {
self.tracked_globally
}
}
impl<T> Drop for MutexGuard<'_, T> {
fn drop(&mut self) {
self.state.owner.store(0, Ordering::Release);
if self.tracked_globally || self.state.contention.has_waiters() {
detector::mutex::release_mutex(self.thread_id, self.lock_id);
} else {
#[cfg(feature = "logging-and-visualization")]
if logger::LOGGING_ENABLED.load(Ordering::Relaxed) {
logger::log_interaction_event(self.thread_id, self.lock_id, Events::MutexReleased);
}
}
}
}
impl<T: Default> Default for Mutex<T> {
fn default() -> Mutex<T> {
Mutex::new(Default::default())
}
}
impl<T> From<T> for Mutex<T> {
fn from(t: T) -> Self {
Mutex::new(t)
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::mem::size_of;
use std::sync::{Arc, mpsc};
use std::time::{Duration, Instant};
#[test]
fn mutex_guard_keeps_one_tracking_reference() {
let maximum_size = size_of::<ParkingLotMutexGuard<'static, ()>>() + 4 * size_of::<usize>();
assert!(
size_of::<MutexGuard<'static, ()>>() <= maximum_size,
"guard stores more than one tracking reference"
);
}
#[test]
fn blocking_mutex_wait_is_visible_until_acquisition() {
let lock = Arc::new(Mutex::new(()));
let owner = lock.lock();
let waiter_lock = Arc::clone(&lock);
let (acquired_tx, acquired_rx) = mpsc::channel();
let waiter = std::thread::spawn(move || {
let _guard = waiter_lock.lock();
acquired_tx.send(()).unwrap();
});
let deadline = Instant::now() + Duration::from_secs(1);
while !lock.state.contention.has_waiters() && Instant::now() < deadline {
std::thread::yield_now();
}
assert!(lock.state.contention.has_waiters());
drop(owner);
acquired_rx.recv_timeout(Duration::from_secs(1)).unwrap();
waiter.join().unwrap();
assert!(!lock.state.contention.has_waiters());
}
}