use core::cell::UnsafeCell;
use core::marker::PhantomData;
use core::ops::{Deref, DerefMut};
use core::sync::atomic::{AtomicUsize, Ordering};
use crate::{Mutex, Semaphore};
struct RawRwLock {
state: AtomicUsize,
mutex: Mutex,
semaphore: Semaphore,
}
const COUNT_BITS: u32 = (usize::BITS - 1) / 2;
const COUNT_MAX: usize = (1usize << COUNT_BITS) - 1;
const IS_WRITING: usize = 1;
const WRITER: usize = 1 << 1;
const READER: usize = 1 << (1 + COUNT_BITS);
const WRITER_MASK: usize = COUNT_MAX << WRITER.trailing_zeros();
const READER_MASK: usize = COUNT_MAX << READER.trailing_zeros();
impl RawRwLock {
const fn new() -> Self {
Self {
state: AtomicUsize::new(0),
mutex: Mutex::new(),
semaphore: Semaphore::new(),
}
}
fn try_lock(&self) -> bool {
if self.mutex.try_lock() {
let state = self.state.load(Ordering::SeqCst);
if state & READER_MASK == 0 {
let _ = self.state.fetch_or(IS_WRITING, Ordering::SeqCst);
return true;
}
self.mutex.unlock();
}
false
}
fn lock(&self) {
let _ = self.state.fetch_add(WRITER, Ordering::SeqCst);
self.mutex.lock();
let state = self
.state
.fetch_add(IS_WRITING.wrapping_sub(WRITER), Ordering::SeqCst);
if state & READER_MASK != 0 {
self.semaphore.wait();
}
}
fn unlock(&self) {
let _ = self.state.fetch_and(!IS_WRITING, Ordering::SeqCst);
self.mutex.unlock();
}
fn try_lock_shared(&self) -> bool {
let state = self.state.load(Ordering::SeqCst);
if state & (IS_WRITING | WRITER_MASK) == 0 {
if self
.state
.compare_exchange(state, state + READER, Ordering::SeqCst, Ordering::SeqCst)
.is_ok()
{
return true;
}
}
if self.mutex.try_lock() {
let _ = self.state.fetch_add(READER, Ordering::SeqCst);
self.mutex.unlock();
return true;
}
false
}
fn lock_shared(&self) {
let mut state = self.state.load(Ordering::SeqCst);
while state & (IS_WRITING | WRITER_MASK) == 0 {
match self.state.compare_exchange_weak(
state,
state + READER,
Ordering::SeqCst,
Ordering::SeqCst,
) {
Ok(_) => return,
Err(s) => state = s,
}
}
self.mutex.lock();
let _ = self.state.fetch_add(READER, Ordering::SeqCst);
self.mutex.unlock();
}
fn unlock_shared(&self) {
let state = self.state.fetch_sub(READER, Ordering::SeqCst);
if (state & READER_MASK == READER) && (state & IS_WRITING != 0) {
self.semaphore.post();
}
}
}
pub struct RwLock<T> {
raw: RawRwLock,
value: UnsafeCell<T>,
}
unsafe impl<T: Send> Send for RwLock<T> {}
unsafe impl<T: Send + Sync> Sync for RwLock<T> {}
impl<T: Default> Default for RwLock<T> {
fn default() -> Self {
Self::new(T::default())
}
}
impl<T> RwLock<T> {
pub const fn new(value: T) -> Self {
Self {
raw: RawRwLock::new(),
value: UnsafeCell::new(value),
}
}
#[inline]
pub fn read(&self) -> RwLockReadGuard<'_, T> {
self.raw.lock_shared();
RwLockReadGuard {
lock: self,
_not_send: PhantomData,
}
}
#[inline]
pub fn write(&self) -> RwLockWriteGuard<'_, T> {
self.raw.lock();
RwLockWriteGuard {
lock: self,
_not_send: PhantomData,
}
}
#[inline]
pub fn try_read(&self) -> Option<RwLockReadGuard<'_, T>> {
if self.raw.try_lock_shared() {
Some(RwLockReadGuard {
lock: self,
_not_send: PhantomData,
})
} else {
None
}
}
#[inline]
pub fn try_write(&self) -> Option<RwLockWriteGuard<'_, T>> {
if self.raw.try_lock() {
Some(RwLockWriteGuard {
lock: self,
_not_send: PhantomData,
})
} else {
None
}
}
#[inline]
pub fn get_mut(&mut self) -> &mut T {
self.value.get_mut()
}
#[inline]
pub fn into_inner(self) -> T {
self.value.into_inner()
}
}
pub struct RwLockReadGuard<'a, T> {
lock: &'a RwLock<T>,
_not_send: PhantomData<*const ()>,
}
impl<'a, T> Deref for RwLockReadGuard<'a, T> {
type Target = T;
#[inline]
fn deref(&self) -> &T {
unsafe { &*self.lock.value.get() }
}
}
impl<'a, T> Drop for RwLockReadGuard<'a, T> {
#[inline]
fn drop(&mut self) {
self.lock.raw.unlock_shared();
}
}
pub struct RwLockWriteGuard<'a, T> {
lock: &'a RwLock<T>,
_not_send: PhantomData<*const ()>,
}
impl<'a, T> Deref for RwLockWriteGuard<'a, T> {
type Target = T;
#[inline]
fn deref(&self) -> &T {
unsafe { &*self.lock.value.get() }
}
}
impl<'a, T> DerefMut for RwLockWriteGuard<'a, T> {
#[inline]
fn deref_mut(&mut self) -> &mut T {
unsafe { &mut *self.lock.value.get() }
}
}
impl<'a, T> Drop for RwLockWriteGuard<'a, T> {
#[inline]
fn drop(&mut self) {
self.lock.raw.unlock();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn smoke() {
let rwl = RwLock::new(0u32);
{
let mut w = rwl.write();
assert!(rwl.try_write().is_none());
assert!(rwl.try_read().is_none());
*w = 1;
}
{
let w = rwl.try_write().unwrap();
assert!(rwl.try_write().is_none());
assert!(rwl.try_read().is_none());
drop(w);
}
{
let r1 = rwl.read();
assert!(rwl.try_write().is_none());
let r2 = rwl.try_read().unwrap();
assert_eq!(*r1, 1);
assert_eq!(*r2, 1);
}
{
let r1 = rwl.try_read().unwrap();
assert!(rwl.try_write().is_none());
let r2 = rwl.try_read().unwrap();
drop((r1, r2));
}
let _w = rwl.write();
}
#[test]
fn raw_internal_state() {
let raw = RawRwLock::new();
raw.lock();
raw.unlock();
assert_eq!(raw.state.load(Ordering::SeqCst), 0);
}
}