moirai_async/sync/
watch.rs1#![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
18pub struct Watch<T> {
20 _phantom: std::marker::PhantomData<T>,
21}
22
23struct WatchState<T> {
24 value: T,
25 version: u64,
26 closed: bool,
27 subscribers: SubscriberRegistry<u64>,
31}
32
33impl<T: Clone + Send + 'static> Watch<T> {
34 #[allow(clippy::new_ret_no_self)] 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
59pub struct WatchSender<T> {
61 state: Arc<Mutex<WatchState<T>>>,
62}
63
64impl<T: Clone> WatchSender<T> {
65 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 pub fn borrow(&self) -> T {
82 self.state.lock().unwrap().value.clone()
83 }
84
85 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 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
120pub struct WatchReceiver<T> {
122 state: Arc<Mutex<WatchState<T>>>,
123 id: u64,
124 version: u64,
125}
126
127impl<T: Clone> WatchReceiver<T> {
128 pub fn borrow(&self) -> T {
130 let state = self.state.lock().unwrap();
131 state.value.clone()
132 }
133
134 pub fn changed(&mut self) -> WatchChanged<'_, T> {
136 WatchChanged { receiver: self }
137 }
138
139 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
176pub 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 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#[derive(Debug, Clone, PartialEq, Eq)]
225pub enum WatchError {
226 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 {}