mod error;
#[cfg(test)]
mod tests;
use std::fmt;
use std::future::Future;
use std::mem;
use std::pin::Pin;
use std::sync::Arc;
use std::task::Context;
use std::task::Poll;
pub use self::error::RecvError;
pub use self::error::SendError;
use crate::internal::mutex::Mutex;
use crate::internal::wake_all;
use crate::internal::wakerset::WakerSet;
use crate::internal::wakerset::WakerToken;
pub fn channel<T: Clone>(initial: T) -> (Sender<T>, Receiver<T>) {
let shared = Arc::new(Shared {
state: Mutex::new(State {
value: initial,
version: 0,
senders: 1,
receivers: 1,
waiters: WakerSet::new(),
}),
});
let sender = Sender {
shared: shared.clone(),
};
let receiver = Receiver { shared, seen: 0 };
(sender, receiver)
}
struct Shared<T> {
state: Mutex<State<T>>,
}
struct State<T> {
value: T,
version: u64,
senders: usize,
receivers: usize,
waiters: WakerSet,
}
pub struct Sender<T> {
shared: Arc<Shared<T>>,
}
impl<T> Clone for Sender<T> {
fn clone(&self) -> Self {
let mut state = self.shared.state.lock();
state.senders = state
.senders
.checked_add(1)
.expect("watch sender count overflowed");
drop(state);
Self {
shared: self.shared.clone(),
}
}
}
impl<T> fmt::Debug for Sender<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Sender").finish_non_exhaustive()
}
}
impl<T> Drop for Sender<T> {
fn drop(&mut self) {
let wakers = {
let mut state = self.shared.state.lock();
state.senders -= 1;
if state.senders != 0 {
return;
}
state.waiters.take_all()
};
wake_all(wakers);
}
}
impl<T> Sender<T> {
pub fn send(&self, value: T) -> Result<(), SendError<T>> {
let (wakers, replaced) = {
let mut state = self.shared.state.lock();
if state.receivers == 0 {
return Err(SendError::new(value));
}
let version = state
.version
.checked_add(1)
.expect("watch channel version counter overflowed");
let replaced = mem::replace(&mut state.value, value);
state.version = version;
let wakers = state.waiters.drain();
(wakers, replaced)
};
wake_all(wakers);
drop(replaced);
Ok(())
}
pub fn send_replace(&self, value: T) -> T {
let (wakers, replaced) = {
let mut state = self.shared.state.lock();
let version = state
.version
.checked_add(1)
.expect("watch channel version counter overflowed");
let replaced = mem::replace(&mut state.value, value);
state.version = version;
let wakers = state.waiters.drain();
(wakers, replaced)
};
wake_all(wakers);
replaced
}
#[must_use = "the receiver is dropped immediately if it is not retained"]
pub fn subscribe(&self) -> Receiver<T> {
let mut state = self.shared.state.lock();
state.receivers = state
.receivers
.checked_add(1)
.expect("watch receiver count overflowed");
let seen = state.version;
drop(state);
Receiver {
shared: self.shared.clone(),
seen,
}
}
}
pub struct Receiver<T> {
shared: Arc<Shared<T>>,
seen: u64,
}
impl<T> Clone for Receiver<T> {
fn clone(&self) -> Self {
let mut state = self.shared.state.lock();
state.receivers = state
.receivers
.checked_add(1)
.expect("watch receiver count overflowed");
drop(state);
Self {
shared: self.shared.clone(),
seen: self.seen,
}
}
}
impl<T> fmt::Debug for Receiver<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("Receiver").finish_non_exhaustive()
}
}
impl<T> Drop for Receiver<T> {
fn drop(&mut self) {
let mut state = self.shared.state.lock();
state.receivers -= 1;
}
}
impl<T> Receiver<T> {
pub fn get(&self) -> T
where
T: Clone,
{
self.shared.state.lock().value.clone()
}
pub fn has_changed(&self) -> Result<bool, RecvError> {
let state = self.shared.state.lock();
if state.version != self.seen {
Ok(true)
} else if state.senders == 0 {
Err(RecvError::Disconnected)
} else {
Ok(false)
}
}
pub async fn changed(&mut self) -> Result<(), RecvError> {
let seen = Change {
shared: &self.shared,
seen: self.seen,
token: None,
}
.await?;
self.seen = seen;
Ok(())
}
pub async fn recv(&mut self) -> Result<T, RecvError>
where
T: Clone,
{
Change {
shared: &self.shared,
seen: self.seen,
token: None,
}
.await?;
let state = self.shared.state.lock();
let value = state.value.clone();
self.seen = state.version;
Ok(value)
}
pub fn is_disconnected(&self) -> bool {
self.shared.state.lock().senders == 0
}
}
struct Change<'a, T> {
shared: &'a Shared<T>,
seen: u64,
token: Option<WakerToken>,
}
impl<T> Future for Change<'_, T> {
type Output = Result<u64, RecvError>;
fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let this = self.get_mut();
let mut state = this.shared.state.lock();
let (poll, retired_waker) = if state.version != this.seen {
this.token = None;
(Poll::Ready(Ok(state.version)), None)
} else if state.senders == 0 {
this.token = None;
(Poll::Ready(Err(RecvError::Disconnected)), None)
} else {
let retired = state.waiters.register(&mut this.token, cx.waker());
(Poll::Pending, retired)
};
drop(state);
drop(retired_waker);
poll
}
}
impl<T> Drop for Change<'_, T> {
fn drop(&mut self) {
if self.token.is_none() {
return;
}
let mut state = self.shared.state.lock();
if state.version != self.seen || state.senders == 0 {
self.token = None;
return;
}
let waker = state.waiters.unregister(&mut self.token);
drop(state);
drop(waker);
}
}