use super::wait_queue::WaitQueue;
use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
use std::sync::{Arc, RwLock};
use std::task::{Context, Poll};
#[repr(align(64))]
struct Shared<T> {
value: RwLock<T>,
version: AtomicU64,
sender_count: AtomicUsize,
receiver_count: AtomicUsize,
wait: WaitQueue,
}
#[must_use]
#[inline]
pub fn channel<T>(init: T) -> (Sender<T>, Receiver<T>) {
let shared = Arc::new(Shared {
value: RwLock::new(init),
version: AtomicU64::new(0),
sender_count: AtomicUsize::new(1),
receiver_count: AtomicUsize::new(1),
wait: WaitQueue::new(),
});
let receiver = Receiver {
shared: shared.clone(),
seen_version: 0,
};
(Sender { shared }, receiver)
}
#[repr(align(64))]
pub struct Sender<T> {
shared: Arc<Shared<T>>,
}
impl<T> Clone for Sender<T> {
#[inline(always)]
fn clone(&self) -> Self {
self.shared.sender_count.fetch_add(1, Ordering::AcqRel);
Self {
shared: self.shared.clone(),
}
}
}
impl<T> Drop for Sender<T> {
#[inline(always)]
fn drop(&mut self) {
if self.shared.sender_count.fetch_sub(1, Ordering::AcqRel) == 1 {
self.shared.wait.wake_all();
}
}
}
impl<T> Sender<T> {
#[inline(always)]
pub fn send(&self, value: T) {
*self
.shared
.value
.write()
.unwrap_or_else(std::sync::PoisonError::into_inner) = value;
self.shared.version.fetch_add(1, Ordering::Release);
self.shared.wait.wake_all();
}
#[inline(always)]
pub fn borrow(&self) -> std::sync::RwLockReadGuard<'_, T> {
self.shared
.value
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
#[must_use]
#[inline(always)]
pub fn is_closed(&self) -> bool {
self.shared.receiver_count.load(Ordering::Acquire) == 0
}
}
#[repr(align(64))]
pub struct Receiver<T> {
shared: Arc<Shared<T>>,
seen_version: u64,
}
impl<T> Clone for Receiver<T> {
#[inline(always)]
fn clone(&self) -> Self {
self.shared.receiver_count.fetch_add(1, Ordering::AcqRel);
Self {
shared: self.shared.clone(),
seen_version: self.seen_version,
}
}
}
impl<T> Drop for Receiver<T> {
#[inline(always)]
fn drop(&mut self) {
self.shared.receiver_count.fetch_sub(1, Ordering::AcqRel);
}
}
impl<T: Send + Sync> Receiver<T> {
#[inline(always)]
pub async fn changed(&mut self) -> Result<(), RecvError> {
std::future::poll_fn(|cx| self.poll_changed(cx)).await
}
}
impl<T> Receiver<T> {
#[inline(always)]
pub fn borrow(&self) -> std::sync::RwLockReadGuard<'_, T> {
self.shared
.value
.read()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
#[inline(always)]
pub fn borrow_and_update(&mut self) -> std::sync::RwLockReadGuard<'_, T> {
self.seen_version = self.shared.version.load(Ordering::Acquire);
self.borrow()
}
#[inline]
fn poll_changed(&mut self, cx: &Context<'_>) -> Poll<Result<(), RecvError>> {
if self.try_observe_change() {
return Poll::Ready(Ok(()));
}
if self.is_closed() {
return Poll::Ready(if self.try_observe_change() {
Ok(())
} else {
Err(RecvError)
});
}
let token = self.shared.wait.register(cx.waker());
if self.try_observe_change() {
self.shared.wait.cancel(token);
return Poll::Ready(Ok(()));
}
if self.is_closed() {
let result = if self.try_observe_change() {
Ok(())
} else {
Err(RecvError)
};
self.shared.wait.cancel(token);
return Poll::Ready(result);
}
Poll::Pending
}
#[inline(always)]
fn try_observe_change(&mut self) -> bool {
let current = self.shared.version.load(Ordering::Acquire);
if current == self.seen_version {
false
} else {
self.seen_version = current;
true
}
}
#[inline(always)]
fn is_closed(&self) -> bool {
self.shared.sender_count.load(Ordering::Acquire) == 0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(align(64))]
pub struct RecvError;
impl std::fmt::Display for RecvError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str("channel closed: every sender dropped")
}
}
impl std::error::Error for RecvError {}