use crate::get_task_from_context;
use crate::runtime::local_executor;
use crate::runtime::task::Task;
use crate::sync::{AsyncRWLock, AsyncReadLockGuard, AsyncWriteLockGuard, LockStatus};
use std::cell::UnsafeCell;
use std::future::Future;
use std::mem::ManuallyDrop;
use std::ops::{Deref, DerefMut};
use std::pin::Pin;
use std::task::{Context, Poll};
pub struct LocalReadLockGuard<'rw_lock, T: ?Sized> {
local_rw_lock: &'rw_lock LocalRWLock<T>,
no_send_marker: std::marker::PhantomData<*const ()>,
}
impl<'rw_lock, T: ?Sized> LocalReadLockGuard<'rw_lock, T> {
#[inline(always)]
fn new(local_rw_lock: &'rw_lock LocalRWLock<T>) -> Self {
Self {
local_rw_lock,
no_send_marker: std::marker::PhantomData,
}
}
}
impl<'rw_lock, T: ?Sized> AsyncReadLockGuard<'rw_lock, T> for LocalReadLockGuard<'rw_lock, T> {
type RWLock = LocalRWLock<T>;
fn rw_lock(&self) -> &'rw_lock Self::RWLock {
self.local_rw_lock
}
#[inline(always)]
unsafe fn leak(self) -> &'rw_lock Self::RWLock {
ManuallyDrop::new(self).local_rw_lock
}
}
impl<T: ?Sized> Deref for LocalReadLockGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.local_rw_lock.get_inner().value
}
}
impl<T: ?Sized> Drop for LocalReadLockGuard<'_, T> {
fn drop(&mut self) {
unsafe {
self.local_rw_lock.read_unlock();
}
}
}
pub struct LocalWriteLockGuard<'rw_lock, T: ?Sized> {
local_rw_lock: &'rw_lock LocalRWLock<T>,
no_send_marker: std::marker::PhantomData<*const ()>,
}
impl<'rw_lock, T: ?Sized> LocalWriteLockGuard<'rw_lock, T> {
#[inline(always)]
fn new(local_rw_lock: &'rw_lock LocalRWLock<T>) -> Self {
Self {
local_rw_lock,
no_send_marker: std::marker::PhantomData,
}
}
}
impl<'rw_lock, T: ?Sized> AsyncWriteLockGuard<'rw_lock, T> for LocalWriteLockGuard<'rw_lock, T> {
type RWLock = LocalRWLock<T>;
fn rw_lock(&self) -> &'rw_lock Self::RWLock {
self.local_rw_lock
}
#[inline(always)]
unsafe fn leak(self) -> &'rw_lock Self::RWLock {
ManuallyDrop::new(self).local_rw_lock
}
}
impl<T: ?Sized> Deref for LocalWriteLockGuard<'_, T> {
type Target = T;
fn deref(&self) -> &Self::Target {
&self.local_rw_lock.get_inner().value
}
}
impl<T: ?Sized> DerefMut for LocalWriteLockGuard<'_, T> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.local_rw_lock.get_inner().value
}
}
impl<T: ?Sized> Drop for LocalWriteLockGuard<'_, T> {
fn drop(&mut self) {
unsafe {
self.local_rw_lock.write_unlock();
}
}
}
pub struct ReadLockWait<'rw_lock, T: ?Sized> {
was_called: bool,
local_rw_lock: &'rw_lock LocalRWLock<T>,
no_send_marker: std::marker::PhantomData<*const ()>,
}
impl<'rw_lock, T: ?Sized> ReadLockWait<'rw_lock, T> {
#[inline(always)]
fn new(local_rw_lock: &'rw_lock LocalRWLock<T>) -> Self {
Self {
was_called: false,
local_rw_lock,
no_send_marker: std::marker::PhantomData,
}
}
}
impl<'rw_lock, T: ?Sized> Future for ReadLockWait<'rw_lock, T> {
type Output = LocalReadLockGuard<'rw_lock, T>;
fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
let this = unsafe { self.get_unchecked_mut() };
if !this.was_called {
let task = unsafe { get_task_from_context!(cx) };
this.local_rw_lock.get_inner().wait_queue_read.push(task);
this.was_called = true;
return Poll::Pending;
}
Poll::Ready(LocalReadLockGuard::new(this.local_rw_lock))
}
}
pub struct WriteLockWait<'rw_lock, T: ?Sized> {
was_called: bool,
local_rw_lock: &'rw_lock LocalRWLock<T>,
no_send_marker: std::marker::PhantomData<*const ()>,
}
impl<'rw_lock, T: ?Sized> WriteLockWait<'rw_lock, T> {
#[inline(always)]
fn new(local_rw_lock: &'rw_lock LocalRWLock<T>) -> Self {
Self {
was_called: false,
local_rw_lock,
no_send_marker: std::marker::PhantomData,
}
}
}
impl<'rw_lock, T: ?Sized> Future for WriteLockWait<'rw_lock, T> {
type Output = LocalWriteLockGuard<'rw_lock, T>;
fn poll(self: Pin<&mut Self>, cx: &mut Context) -> Poll<Self::Output> {
let this = unsafe { self.get_unchecked_mut() };
if !this.was_called {
let task = unsafe { get_task_from_context!(cx) };
this.local_rw_lock.get_inner().wait_queue_write.push(task);
this.was_called = true;
return Poll::Pending;
}
Poll::Ready(LocalWriteLockGuard::new(this.local_rw_lock))
}
}
struct Inner<T: ?Sized> {
wait_queue_read: Vec<Task>,
wait_queue_write: Vec<Task>,
number_of_readers: isize,
value: T,
}
pub struct LocalRWLock<T: ?Sized> {
no_send_marker: std::marker::PhantomData<*const ()>,
inner: UnsafeCell<Inner<T>>,
}
impl<T: ?Sized> LocalRWLock<T> {
#[inline(always)]
pub const fn new(value: T) -> Self
where
T: Sized,
{
Self {
inner: UnsafeCell::new(Inner {
wait_queue_read: Vec::new(),
wait_queue_write: Vec::new(),
number_of_readers: 0,
value,
}),
no_send_marker: std::marker::PhantomData,
}
}
#[inline(always)]
#[allow(clippy::mut_from_ref, reason = "It is Sync and `local`")]
fn get_inner(&self) -> &mut Inner<T> {
unsafe { &mut *self.inner.get() }
}
}
impl<T: ?Sized> AsyncRWLock<T> for LocalRWLock<T> {
type ReadLockGuard<'rw_lock>
= LocalReadLockGuard<'rw_lock, T>
where
T: 'rw_lock,
Self: 'rw_lock;
type WriteLockGuard<'rw_lock>
= LocalWriteLockGuard<'rw_lock, T>
where
T: 'rw_lock,
Self: 'rw_lock;
#[inline(always)]
fn get_lock_status(&self) -> LockStatus {
#[allow(clippy::cast_sign_loss, reason = "false positive")]
match self.get_inner().number_of_readers {
0 => LockStatus::Unlocked,
n if n > 0 => LockStatus::ReadLocked(n as usize),
_ => LockStatus::WriteLocked,
}
}
#[inline(always)]
#[allow(clippy::future_not_send, reason = "Because it is `local`")]
async fn write<'rw_lock>(&'rw_lock self) -> Self::WriteLockGuard<'rw_lock>
where
T: 'rw_lock,
{
let inner = self.get_inner();
if inner.number_of_readers == 0 {
debug_assert!(inner.wait_queue_read.is_empty());
inner.number_of_readers = -1;
return LocalWriteLockGuard::new(self);
}
WriteLockWait::new(self).await
}
#[inline(always)]
#[allow(clippy::future_not_send, reason = "Because it is `local`")]
async fn read<'rw_lock>(&'rw_lock self) -> Self::ReadLockGuard<'rw_lock>
where
T: 'rw_lock,
{
let inner = self.get_inner();
if inner.number_of_readers > -1 {
inner.number_of_readers += 1;
return LocalReadLockGuard::new(self);
}
ReadLockWait::new(self).await
}
#[inline(always)]
fn try_write(&self) -> Option<Self::WriteLockGuard<'_>> {
let inner = self.get_inner();
if inner.number_of_readers == 0 {
debug_assert!(inner.wait_queue_read.is_empty());
inner.number_of_readers = -1;
Some(LocalWriteLockGuard::new(self))
} else {
None
}
}
#[inline(always)]
fn try_read(&self) -> Option<Self::ReadLockGuard<'_>> {
let inner = self.get_inner();
if inner.number_of_readers > -1 {
inner.number_of_readers += 1;
Some(LocalReadLockGuard::new(self))
} else {
None
}
}
#[inline(always)]
fn get_mut(&mut self) -> &mut T {
&mut self.inner.get_mut().value
}
#[inline(always)]
unsafe fn read_unlock(&self) {
if cfg!(debug_assertions) {
assert_ne!(
self.get_inner().number_of_readers,
-1,
"LocalRWLock is locked for write"
);
assert_ne!(
self.get_inner().number_of_readers,
0,
"LocalRWLock is already unlocked"
);
}
let inner = self.get_inner();
inner.number_of_readers -= 1;
if inner.number_of_readers == 0 {
debug_assert!(inner.wait_queue_read.is_empty());
let task = inner.wait_queue_write.pop();
if task.is_some() {
inner.number_of_readers = -1;
local_executor().exec_task(unsafe { task.unwrap_unchecked() });
}
}
}
#[inline(always)]
unsafe fn write_unlock(&self) {
if cfg!(debug_assertions) {
assert_ne!(
self.get_inner().number_of_readers,
0,
"LocalRWLock is already unlocked"
);
assert!(
self.get_inner().number_of_readers <= 0,
"LocalRWLock is locked for read"
);
}
let inner = self.get_inner();
let task = inner.wait_queue_write.pop();
if task.is_none() {
let mut readers_count = inner.wait_queue_read.len();
#[allow(clippy::cast_possible_wrap, reason = "false positive")]
{
inner.number_of_readers = readers_count as isize;
}
while readers_count > 0 {
let task = inner.wait_queue_read.pop();
local_executor().exec_task(unsafe { task.unwrap_unchecked() });
readers_count -= 1;
}
} else {
local_executor().exec_task(unsafe { task.unwrap_unchecked() });
}
}
#[inline(always)]
unsafe fn get_read_locked(&self) -> Self::ReadLockGuard<'_> {
if cfg!(debug_assertions) {
assert_ne!(
self.get_inner().number_of_readers,
-1,
"LocalRWLock is locked for write"
);
assert_ne!(
self.get_inner().number_of_readers,
0,
"LocalRWLock is unlocked"
);
}
LocalReadLockGuard::new(self)
}
#[inline(always)]
unsafe fn get_write_locked(&self) -> Self::WriteLockGuard<'_> {
if cfg!(debug_assertions) {
assert_ne!(
self.get_inner().number_of_readers,
0,
"LocalRWLock is unlocked, but get_write_locked is called"
);
assert!(
self.get_inner().number_of_readers <= 0,
"LocalRWLock is locked for read"
);
}
LocalWriteLockGuard::new(self)
}
}
unsafe impl<T: ?Sized + Sync> Sync for LocalRWLock<T> {}
#[allow(dead_code, reason = "It is used only in compile tests")]
fn test_compile_local_rw_lock() {}
#[cfg(test)]
mod tests {
use super::*;
use crate as orengine;
use crate::sync::{AsyncWaitGroup, LocalWaitGroup};
use crate::yield_now;
use std::rc::Rc;
#[orengine::test::test_local]
fn test_local_rw_lock() {
let rw_lock = Rc::new(LocalRWLock::new(0));
let wg = Rc::new(LocalWaitGroup::new());
let read_wg = Rc::new(LocalWaitGroup::new());
for i in 1..=15 {
let mutex = rw_lock.clone();
local_executor().exec_local_future(async move {
let value = mutex.read().await;
assert_eq!(mutex.get_inner().number_of_readers, i);
assert_eq!(*value, 0);
yield_now().await;
assert_eq!(mutex.get_inner().number_of_readers, 16 - i);
assert_eq!(*value, 0);
});
}
for _ in 1..=15 {
let wg = wg.clone();
let read_wg = read_wg.clone();
wg.add(1);
let mutex = rw_lock.clone();
local_executor().exec_local_future(async move {
assert_eq!(mutex.get_inner().number_of_readers, 15);
let mut value = mutex.write().await;
{
let read_wg = read_wg.clone();
let mutex = mutex.clone();
read_wg.add(1);
local_executor().exec_local_future(async move {
assert_eq!(mutex.get_inner().number_of_readers, -1);
let value = mutex.read().await;
assert_ne!(*value, 0);
assert_ne!(mutex.get_inner().number_of_readers, 0);
read_wg.done();
});
}
assert_eq!(mutex.get_inner().number_of_readers, -1);
*value += 1;
wg.done();
});
}
wg.wait().await;
read_wg.wait().await;
let value = rw_lock.read().await;
assert_eq!(*value, 15);
assert_ne!(rw_lock.get_inner().number_of_readers, 0);
}
#[orengine::test::test_local]
fn test_try_local_rw_lock() {
const NUMBER_OF_READERS: isize = 5;
let rw_lock = Rc::new(LocalRWLock::new(0));
for i in 1..=NUMBER_OF_READERS {
let mutex = rw_lock.clone();
local_executor().exec_local_future(async move {
let lock = mutex.try_read().expect("Failed to get read lock!");
assert_eq!(mutex.get_inner().number_of_readers, i);
yield_now().await;
drop(lock);
});
}
assert_eq!(rw_lock.get_inner().number_of_readers, NUMBER_OF_READERS);
assert!(
rw_lock.try_write().is_none(),
"Successful attempt to acquire write lock when rw_lock locked for read"
);
yield_now().await;
assert_eq!(rw_lock.get_inner().number_of_readers, 0);
let mut write_lock = rw_lock.try_write().expect("Failed to get write lock!");
*write_lock += 1;
assert_eq!(*write_lock, 1);
assert_eq!(rw_lock.get_inner().number_of_readers, -1);
assert!(
rw_lock.try_read().is_none(),
"Successful attempt to acquire read lock when rw_lock locked for write"
);
assert!(
rw_lock.try_write().is_none(),
"Successful attempt to acquire write lock when rw_lock locked for write"
);
}
}