use moirai_utils::CacheAligned;
use std::cell::UnsafeCell;
use std::fmt;
use std::hint;
use std::ops::{Deref, DerefMut};
use std::sync::atomic::{AtomicBool, Ordering};
const SPINLOCK_MAX_BACKOFF: usize = 64;
const SPINLOCK_MAX_SPINS_BEFORE_YIELD: usize = 1000;
const SPINLOCK_INITIAL_BACKOFF: usize = 1;
pub struct SpinLock<T> {
locked: CacheAligned<AtomicBool>,
data: UnsafeCell<T>,
}
impl<T> fmt::Debug for SpinLock<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let locked = self.locked.load(Ordering::Relaxed);
f.debug_struct("SpinLock")
.field("locked", &locked)
.finish_non_exhaustive()
}
}
unsafe impl<T: Send> Send for SpinLock<T> {}
unsafe impl<T: Send> Sync for SpinLock<T> {}
impl<T> SpinLock<T> {
pub const fn new(data: T) -> Self {
Self {
locked: CacheAligned::new(AtomicBool::new(false)),
data: UnsafeCell::new(data),
}
}
pub fn lock(&self) -> SpinLockGuard<'_, T> {
let mut backoff = SPINLOCK_INITIAL_BACKOFF;
let mut total_spins = 0;
loop {
if !self.locked.load(Ordering::Relaxed)
&& self
.locked
.compare_exchange_weak(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
{
return SpinLockGuard {
lock: self,
_phantom: std::marker::PhantomData,
};
}
for _ in 0..backoff {
hint::spin_loop();
}
if backoff < SPINLOCK_MAX_BACKOFF {
backoff = backoff.saturating_mul(2);
}
total_spins += backoff;
if total_spins >= SPINLOCK_MAX_SPINS_BEFORE_YIELD {
std::thread::yield_now();
total_spins = 0;
backoff = SPINLOCK_INITIAL_BACKOFF; }
}
}
pub fn try_lock(&self) -> Option<SpinLockGuard<'_, T>> {
if !self.locked.load(Ordering::Relaxed)
&& self
.locked
.compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
{
Some(SpinLockGuard {
lock: self,
_phantom: std::marker::PhantomData,
})
} else {
None
}
}
}
pub struct SpinLockGuard<'a, T> {
lock: &'a SpinLock<T>,
_phantom: std::marker::PhantomData<T>,
}
impl<'a, T> Drop for SpinLockGuard<'a, T> {
fn drop(&mut self) {
self.lock.locked.store(false, Ordering::Release);
}
}
impl<'a, T> Deref for SpinLockGuard<'a, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
unsafe { &*self.lock.data.get() }
}
}
impl<'a, T> DerefMut for SpinLockGuard<'a, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
unsafe { &mut *self.lock.data.get() }
}
}