use std::collections::BTreeMap;
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll, Waker};
pub struct Watch<T> {
_phantom: std::marker::PhantomData<T>,
}
struct WatchState<T> {
value: T,
version: u64,
closed: bool,
receivers: BTreeMap<u64, WatchReceiverState>,
next_receiver_id: u64,
}
struct WatchReceiverState {
version: u64,
waker: Option<Waker>,
}
impl<T: Clone + Send + 'static> Watch<T> {
#[allow(clippy::new_ret_no_self)] pub fn new(initial: T) -> (WatchSender<T>, WatchReceiver<T>) {
let state = Arc::new(Mutex::new(WatchState {
value: initial,
version: 0,
closed: false,
receivers: BTreeMap::new(),
next_receiver_id: 1,
}));
let sender = WatchSender {
state: state.clone(),
};
let receiver = WatchReceiver {
state: state.clone(),
id: 0,
version: 0,
};
{
let mut state_guard = state.lock().unwrap();
state_guard.receivers.insert(
0,
WatchReceiverState {
version: 0,
waker: None,
},
);
}
(sender, receiver)
}
}
pub struct WatchSender<T> {
state: Arc<Mutex<WatchState<T>>>,
}
impl<T: Clone> WatchSender<T> {
pub fn send(&self, value: T) -> Result<(), WatchError> {
let mut state = self.state.lock().unwrap();
if state.closed {
return Err(WatchError::Closed);
}
state.value = value;
state.version += 1;
let current_version = state.version;
for receiver in state.receivers.values_mut() {
if receiver.version < current_version {
if let Some(waker) = receiver.waker.take() {
waker.wake();
}
}
}
Ok(())
}
pub fn borrow(&self) -> T {
self.state.lock().unwrap().value.clone()
}
pub fn send_modify<F>(&self, modify: F) -> Result<(), WatchError>
where
F: FnOnce(&mut T),
{
let mut state = self.state.lock().unwrap();
if state.closed {
return Err(WatchError::Closed);
}
modify(&mut state.value);
state.version += 1;
for receiver in state.receivers.values_mut() {
if let Some(waker) = receiver.waker.take() {
waker.wake();
}
}
Ok(())
}
pub fn receiver_count(&self) -> usize {
self.state.lock().unwrap().receivers.len()
}
}
impl<T> Drop for WatchSender<T> {
fn drop(&mut self) {
let mut state = self.state.lock().unwrap();
state.closed = true;
for receiver in state.receivers.values_mut() {
if let Some(waker) = receiver.waker.take() {
waker.wake();
}
}
}
}
pub struct WatchReceiver<T> {
state: Arc<Mutex<WatchState<T>>>,
id: u64,
version: u64,
}
impl<T: Clone> WatchReceiver<T> {
pub fn borrow(&self) -> T {
let state = self.state.lock().unwrap();
state.value.clone()
}
pub fn changed(&mut self) -> WatchChanged<'_, T> {
WatchChanged { receiver: self }
}
pub fn has_changed(&mut self) -> bool {
let mut state = self.state.lock().unwrap();
let changed = state.version > self.version;
if changed {
let current_version = state.version;
self.version = current_version;
if let Some(receiver_state) = state.receivers.get_mut(&self.id) {
receiver_state.version = current_version;
}
}
changed
}
}
impl<T> Clone for WatchReceiver<T> {
fn clone(&self) -> Self {
let mut state = self.state.lock().unwrap();
let new_id = state.next_receiver_id;
state.next_receiver_id += 1;
let current_version = state.version;
state.receivers.insert(
new_id,
WatchReceiverState {
version: current_version,
waker: None,
},
);
WatchReceiver {
state: self.state.clone(),
id: new_id,
version: current_version,
}
}
}
impl<T> Drop for WatchReceiver<T> {
fn drop(&mut self) {
if let Ok(mut state) = self.state.lock() {
state.receivers.remove(&self.id);
}
}
}
pub struct WatchChanged<'a, T> {
receiver: &'a mut WatchReceiver<T>,
}
impl<'a, T: Clone> Future for WatchChanged<'a, T> {
type Output = Result<(), WatchError>;
fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
let receiver = &mut *self.receiver;
let mut state = receiver.state.lock().unwrap();
if state.closed {
return Poll::Ready(Err(WatchError::Closed));
}
let current_version = state.version;
if current_version > receiver.version {
receiver.version = current_version;
if let Some(receiver_state) = state.receivers.get_mut(&receiver.id) {
receiver_state.version = current_version;
}
return Poll::Ready(Ok(()));
}
if let Some(receiver_state) = state.receivers.get_mut(&receiver.id) {
receiver_state.waker = Some(cx.waker().clone());
}
Poll::Pending
}
}
impl<'a, T> Drop for WatchChanged<'a, T> {
fn drop(&mut self) {
if let Ok(mut state) = self.receiver.state.lock() {
if let Some(receiver_state) = state.receivers.get_mut(&self.receiver.id) {
receiver_state.waker = None;
}
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum WatchError {
Closed,
}
impl std::fmt::Display for WatchError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
WatchError::Closed => write!(f, "watch channel is closed"),
}
}
}
impl std::error::Error for WatchError {}