use core::{
cell::UnsafeCell,
hint::spin_loop,
marker::PhantomData,
mem::ManuallyDrop,
ops::{Deref, DerefMut},
ptr::NonNull,
sync::atomic::{AtomicUsize, Ordering},
};
#[derive(Debug)]
pub(crate) struct RawTicketLock<T> {
next: AtomicUsize,
owner: AtomicUsize,
value: UnsafeCell<T>,
}
impl<T> RawTicketLock<T> {
pub(crate) const fn new(value: T) -> Self {
Self {
next: AtomicUsize::new(0),
owner: AtomicUsize::new(0),
value: UnsafeCell::new(value),
}
}
pub(crate) fn lock(&self) -> RawTicketGuard<'_, T> {
let ticket = self.next.fetch_add(1, Ordering::Relaxed);
while self.owner.load(Ordering::Acquire) != ticket {
spin_loop();
}
RawTicketGuard {
lock: self,
ticket,
_not_send: PhantomData,
}
}
pub(crate) fn try_lock(&self) -> Option<RawTicketGuard<'_, T>> {
let owner = self.owner.load(Ordering::Acquire);
self.next
.compare_exchange(
owner,
owner.wrapping_add(1),
Ordering::Acquire,
Ordering::Relaxed,
)
.ok()
.map(|ticket| RawTicketGuard {
lock: self,
ticket,
_not_send: PhantomData,
})
}
fn unlock(&self, ticket: usize) {
self.owner.store(ticket.wrapping_add(1), Ordering::Release);
}
}
unsafe impl<T: Send> Send for RawTicketLock<T> {}
unsafe impl<T: Send> Sync for RawTicketLock<T> {}
pub(crate) struct RawTicketGuard<'a, T> {
lock: &'a RawTicketLock<T>,
ticket: usize,
_not_send: PhantomData<*mut ()>,
}
#[derive(Debug)]
pub(crate) struct RawTicketBaton<T> {
lock: NonNull<RawTicketLock<T>>,
ticket: usize,
_not_send: PhantomData<*mut T>,
}
impl<T> RawTicketGuard<'_, T> {
pub(crate) unsafe fn into_baton(self) -> RawTicketBaton<T> {
let this = ManuallyDrop::new(self);
RawTicketBaton {
lock: NonNull::from(this.lock),
ticket: this.ticket,
_not_send: PhantomData,
}
}
}
impl<T> Deref for RawTicketGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
unsafe { &*self.lock.value.get() }
}
}
impl<T> DerefMut for RawTicketGuard<'_, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
unsafe { &mut *self.lock.value.get() }
}
}
impl<T> Drop for RawTicketGuard<'_, T> {
fn drop(&mut self) {
self.lock.unlock(self.ticket);
}
}
impl<T> Drop for RawTicketBaton<T> {
fn drop(&mut self) {
unsafe { self.lock.as_ref() }.unlock(self.ticket);
}
}
#[cfg(test)]
mod tests {
use alloc::sync::Arc;
use super::*;
#[test]
fn try_lock_does_not_consume_a_ticket_on_failure() {
let lock = RawTicketLock::new(0usize);
let first = lock.lock();
assert!(lock.try_lock().is_none());
drop(first);
let mut second = lock.try_lock().expect("failed try-lock must roll back");
*second = 1;
assert_eq!(*second, 1);
}
#[test]
fn serializes_concurrent_writers() {
let lock = Arc::new(RawTicketLock::new(0usize));
let workers: alloc::vec::Vec<_> = (0..4)
.map(|_| {
let lock = Arc::clone(&lock);
std::thread::spawn(move || {
for _ in 0..1_000 {
*lock.lock() += 1;
}
})
})
.collect();
for worker in workers {
worker.join().expect("writer thread panicked");
}
assert_eq!(*lock.lock(), 4_000);
}
}