Skip to main content

moirai_async/sync/
watch.rs

1//! Watch channel for state monitoring with change notifications
2//!
3//! Provides watch channel implementation that allows monitoring state changes
4//! with async notifications, following SLAP principle design.
5
6#![expect(
7    clippy::unwrap_used,
8    reason = "ratchet MOIRAI-UNWRAP-1: pre-existing debt"
9)]
10
11use std::future::Future;
12use std::pin::Pin;
13use std::sync::{Arc, Mutex};
14use std::task::{Context, Poll};
15
16use super::subscribers::{SubscriberRegistry, wake_drained};
17
18/// Watch channel for state monitoring with change notifications
19pub struct Watch<T> {
20    _phantom: std::marker::PhantomData<T>,
21}
22
23struct WatchState<T> {
24    value: T,
25    version: u64,
26    closed: bool,
27    /// One slot per receiver, shared with `Broadcast` (see `subscribers`). The
28    /// cursor is the version that receiver last observed; the slot's waker is
29    /// registered while it waits for a newer one.
30    subscribers: SubscriberRegistry<u64>,
31}
32
33impl<T: Clone + Send + 'static> Watch<T> {
34    /// Create a new watch channel with an initial value
35    /// Returns (sender, receiver) tuple per channel pattern conventions
36    #[allow(clippy::new_ret_no_self)] // Standard channel pattern per Rust Book Ch.16
37    pub fn new(initial: T) -> (WatchSender<T>, WatchReceiver<T>) {
38        let state = Arc::new(Mutex::new(WatchState {
39            value: initial,
40            version: 0,
41            closed: false,
42            subscribers: SubscriberRegistry::with_initial(0),
43        }));
44
45        let sender = WatchSender {
46            state: state.clone(),
47        };
48
49        let receiver = WatchReceiver {
50            state: state.clone(),
51            id: 0,
52            version: 0,
53        };
54
55        (sender, receiver)
56    }
57}
58
59/// Sender half of watch channel
60pub struct WatchSender<T> {
61    state: Arc<Mutex<WatchState<T>>>,
62}
63
64impl<T: Clone> WatchSender<T> {
65    /// Send a new value, notifying all receivers
66    pub fn send(&self, value: T) -> Result<(), WatchError> {
67        let wakers = {
68            let mut state = self.state.lock().unwrap();
69            if state.closed {
70                return Err(WatchError::Closed);
71            }
72            state.value = value;
73            state.version += 1;
74            state.subscribers.drain_wakers()
75        };
76        wake_drained(wakers);
77        Ok(())
78    }
79
80    /// Get the current value
81    pub fn borrow(&self) -> T {
82        self.state.lock().unwrap().value.clone()
83    }
84
85    /// Modify the value in place and notify receivers
86    pub fn send_modify<F>(&self, modify: F) -> Result<(), WatchError>
87    where
88        F: FnOnce(&mut T),
89    {
90        let wakers = {
91            let mut state = self.state.lock().unwrap();
92            if state.closed {
93                return Err(WatchError::Closed);
94            }
95            modify(&mut state.value);
96            state.version += 1;
97            state.subscribers.drain_wakers()
98        };
99        wake_drained(wakers);
100        Ok(())
101    }
102
103    /// Get the number of active receivers
104    pub fn receiver_count(&self) -> usize {
105        self.state.lock().unwrap().subscribers.len()
106    }
107}
108
109impl<T> Drop for WatchSender<T> {
110    fn drop(&mut self) {
111        let wakers = {
112            let mut state = self.state.lock().unwrap();
113            state.closed = true;
114            state.subscribers.drain_wakers()
115        };
116        wake_drained(wakers);
117    }
118}
119
120/// Receiver half of watch channel
121pub struct WatchReceiver<T> {
122    state: Arc<Mutex<WatchState<T>>>,
123    id: u64,
124    version: u64,
125}
126
127impl<T: Clone> WatchReceiver<T> {
128    /// Get the current value
129    pub fn borrow(&self) -> T {
130        let state = self.state.lock().unwrap();
131        state.value.clone()
132    }
133
134    /// Wait for the value to change
135    pub fn changed(&mut self) -> WatchChanged<'_, T> {
136        WatchChanged { receiver: self }
137    }
138
139    /// Check if the value has changed since last check
140    pub fn has_changed(&mut self) -> bool {
141        let mut state = self.state.lock().unwrap();
142        let changed = state.version > self.version;
143        if changed {
144            let current_version = state.version;
145            self.version = current_version;
146            if let Some(subscriber) = state.subscribers.get_mut(self.id) {
147                subscriber.cursor = current_version;
148            }
149        }
150        changed
151    }
152}
153
154impl<T> Clone for WatchReceiver<T> {
155    fn clone(&self) -> Self {
156        let mut state = self.state.lock().unwrap();
157        let current_version = state.version;
158        let id = state.subscribers.register(current_version);
159
160        WatchReceiver {
161            state: self.state.clone(),
162            id,
163            version: current_version,
164        }
165    }
166}
167
168impl<T> Drop for WatchReceiver<T> {
169    fn drop(&mut self) {
170        if let Ok(mut state) = self.state.lock() {
171            state.subscribers.remove(self.id);
172        }
173    }
174}
175
176/// Future for waiting for watch value changes
177pub struct WatchChanged<'a, T> {
178    receiver: &'a mut WatchReceiver<T>,
179}
180
181impl<'a, T: Clone> Future for WatchChanged<'a, T> {
182    type Output = Result<(), WatchError>;
183
184    fn poll(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
185        let receiver = &mut *self.receiver;
186        let mut state = receiver.state.lock().unwrap();
187
188        if state.closed {
189            return Poll::Ready(Err(WatchError::Closed));
190        }
191
192        let current_version = state.version;
193        if current_version > receiver.version {
194            receiver.version = current_version;
195            if let Some(subscriber) = state.subscribers.get_mut(receiver.id) {
196                subscriber.cursor = current_version;
197            }
198            return Poll::Ready(Ok(()));
199        }
200
201        if let Some(subscriber) = state.subscribers.get_mut(receiver.id) {
202            subscriber.waker = Some(cx.waker().clone());
203        }
204
205        Poll::Pending
206    }
207}
208
209impl<'a, T> Drop for WatchChanged<'a, T> {
210    fn drop(&mut self) {
211        // If this future is dropped while pending, the waker left in the
212        // subscriber slot would be called by the next `send()` on a
213        // now-deallocated task allocation — a use-after-free of the waker.
214        // Clear it here so the sender only wakes live futures.
215        if let Ok(mut state) = self.receiver.state.lock()
216            && let Some(subscriber) = state.subscribers.get_mut(self.receiver.id)
217        {
218            subscriber.waker = None;
219        }
220    }
221}
222
223/// Error types for watch channel operations
224#[derive(Debug, Clone, PartialEq, Eq)]
225pub enum WatchError {
226    /// Channel has been closed
227    Closed,
228}
229
230impl std::fmt::Display for WatchError {
231    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
232        match self {
233            WatchError::Closed => write!(f, "watch channel is closed"),
234        }
235    }
236}
237
238impl std::error::Error for WatchError {}