#![expect(
clippy::unwrap_used,
reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
)]
use std::future::Future;
use std::pin::Pin;
use std::sync::{Arc, Mutex};
use std::task::{Context, Poll};
use super::subscribers::{SubscriberRegistry, wake_drained};
pub struct Watch<T> {
_phantom: std::marker::PhantomData<T>,
}
struct WatchState<T> {
value: T,
version: u64,
closed: bool,
subscribers: SubscriberRegistry<u64>,
}
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,
subscribers: SubscriberRegistry::with_initial(0),
}));
let sender = WatchSender {
state: state.clone(),
};
let receiver = WatchReceiver {
state: state.clone(),
id: 0,
version: 0,
};
(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 wakers = {
let mut state = self.state.lock().unwrap();
if state.closed {
return Err(WatchError::Closed);
}
state.value = value;
state.version += 1;
state.subscribers.drain_wakers()
};
wake_drained(wakers);
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 wakers = {
let mut state = self.state.lock().unwrap();
if state.closed {
return Err(WatchError::Closed);
}
modify(&mut state.value);
state.version += 1;
state.subscribers.drain_wakers()
};
wake_drained(wakers);
Ok(())
}
pub fn receiver_count(&self) -> usize {
self.state.lock().unwrap().subscribers.len()
}
}
impl<T> Drop for WatchSender<T> {
fn drop(&mut self) {
let wakers = {
let mut state = self.state.lock().unwrap();
state.closed = true;
state.subscribers.drain_wakers()
};
wake_drained(wakers);
}
}
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(subscriber) = state.subscribers.get_mut(self.id) {
subscriber.cursor = current_version;
}
}
changed
}
}
impl<T> Clone for WatchReceiver<T> {
fn clone(&self) -> Self {
let mut state = self.state.lock().unwrap();
let current_version = state.version;
let id = state.subscribers.register(current_version);
WatchReceiver {
state: self.state.clone(),
id,
version: current_version,
}
}
}
impl<T> Drop for WatchReceiver<T> {
fn drop(&mut self) {
if let Ok(mut state) = self.state.lock() {
state.subscribers.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(subscriber) = state.subscribers.get_mut(receiver.id) {
subscriber.cursor = current_version;
}
return Poll::Ready(Ok(()));
}
if let Some(subscriber) = state.subscribers.get_mut(receiver.id) {
subscriber.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()
&& let Some(subscriber) = state.subscribers.get_mut(self.receiver.id)
{
subscriber.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 {}