#![cfg_attr(all(doc, not(doctest)), doc = include_str!("../README.md"))]
#![cfg_attr(
any(not(doc), doctest),
doc = "Async single-producer, multi-consumer channel that only retains the last sent value"
)]
use {
event_listener::Event,
futures_lite::{Stream, StreamExt, stream},
std::{
error, fmt,
pin::Pin,
sync::{
Arc, RwLock, RwLockReadGuard, RwLockWriteGuard,
atomic::{AtomicUsize, Ordering},
},
task::{Context, Poll},
},
};
pub fn channel<T>(init: T) -> (Sender<T>, Receiver<T>) {
let shared = Arc::new(Shared {
value: RwLock::new(init),
state: State::new(),
rx_count: AtomicUsize::new(1),
changed: Event::new(),
all_receivers_dropped: Event::new(),
});
let tx = Sender {
shared: shared.clone(),
};
let rx = Receiver {
shared,
last_version: 0,
};
(tx, rx)
}
#[derive(Debug)]
pub struct Sender<T> {
shared: Arc<Shared<T>>,
}
impl<T> Sender<T> {
pub fn send(&self, value: T) -> Result<(), SendError<T>> {
if self.shared.rx_count.load(Ordering::Relaxed) == 0 {
return Err(SendError(value));
}
*self.shared.write_value() = value;
self.shared.state.increment_version();
self.shared.changed.notify(usize::MAX);
Ok(())
}
pub async fn closed(&self) {
if self.shared.rx_count.load(Ordering::Relaxed) == 0 {
return;
}
event_listener::listener!(self.shared.all_receivers_dropped => listener);
if self.shared.rx_count.load(Ordering::Relaxed) == 0 {
return;
}
listener.await;
debug_assert_eq!(self.shared.rx_count.load(Ordering::Relaxed), 0);
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
self.shared.state.close();
self.shared.changed.notify(usize::MAX);
}
}
#[derive(PartialEq, Eq)]
pub struct SendError<T>(pub T);
impl<T> fmt::Display for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("sending on a closed channel")
}
}
impl<T> fmt::Debug for SendError<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("sending on a closed channel")
}
}
impl<T> error::Error for SendError<T> {}
#[derive(Debug)]
pub struct Receiver<T> {
shared: Arc<Shared<T>>,
last_version: usize,
}
impl<T> Receiver<T> {
pub fn observe<F, R>(&self, f: F) -> R
where
F: FnOnce(&T) -> R,
{
f(&self.shared.read_value())
}
pub async fn changed(&mut self) -> Result<(), RecvError> {
if self
.shared
.state
.version_changed(&mut self.last_version)
.ok_or(RecvError)?
{
return Ok(());
}
event_listener::listener!(self.shared.changed => listener);
if self
.shared
.state
.version_changed(&mut self.last_version)
.ok_or(RecvError)?
{
return Ok(());
}
listener.await;
let changed = self
.shared
.state
.version_changed(&mut self.last_version)
.ok_or(RecvError)?;
debug_assert!(changed);
Ok(())
}
pub async fn recv(&mut self) -> Result<T, RecvError>
where
T: Clone,
{
self.changed().await?;
Ok(self.observe(T::clone))
}
pub fn into_stream<'item>(self) -> Updates<'item, T>
where
T: Clone + Send + Sync + 'item,
{
Updates {
inner: Box::pin(stream::unfold(self, async |mut me| {
let value = me.recv().await.ok()?;
Some((value, me))
})),
}
}
}
impl<T> Clone for Receiver<T> {
fn clone(&self) -> Self {
self.shared.rx_count.fetch_add(1, Ordering::Relaxed);
Self {
shared: self.shared.clone(),
last_version: self.last_version,
}
}
}
impl<T> Drop for Receiver<T> {
fn drop(&mut self) {
if self.shared.rx_count.fetch_sub(1, Ordering::Relaxed) == 1 {
self.shared.all_receivers_dropped.notify(usize::MAX);
}
}
}
#[derive(PartialEq, Eq)]
pub struct RecvError;
impl fmt::Display for RecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("receiving on a closed channel")
}
}
impl fmt::Debug for RecvError {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str("receiving on a closed channel")
}
}
impl error::Error for RecvError {}
pub struct Updates<'item, T> {
inner: Pin<Box<dyn Stream<Item = T> + Send + Sync + 'item>>,
}
impl<T> Stream for Updates<'_, T> {
type Item = T;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.inner.poll_next(cx)
}
}
#[derive(Debug)]
struct Shared<T> {
value: RwLock<T>,
state: State,
rx_count: AtomicUsize,
changed: Event,
all_receivers_dropped: Event,
}
impl<T> Shared<T> {
fn read_value(&self) -> RwLockReadGuard<'_, T> {
match self.value.read() {
Ok(guard) => guard,
Err(e) => e.into_inner(),
}
}
fn write_value(&self) -> RwLockWriteGuard<'_, T> {
match self.value.write() {
Ok(guard) => guard,
Err(e) => e.into_inner(),
}
}
}
#[derive(Debug)]
struct State(AtomicUsize);
impl State {
const VERSION_STEP: usize = 2;
const CLOSED_BIT: usize = 1;
fn new() -> Self {
Self(AtomicUsize::new(0))
}
fn increment_version(&self) {
self.0.fetch_add(Self::VERSION_STEP, Ordering::Release);
}
fn version_changed(&self, last_version: &mut usize) -> Option<bool> {
let state = self.0.load(Ordering::Acquire);
let new_version = state & !Self::CLOSED_BIT;
if *last_version != new_version {
*last_version = new_version;
return Some(true);
}
if Self::CLOSED_BIT == state & Self::CLOSED_BIT {
return None;
}
Some(false)
}
fn close(&self) {
self.0.fetch_or(Self::CLOSED_BIT, Ordering::Release);
}
}