Skip to main content

dtact_util/sync/
watch.rs

1//! Single-value "latest state" broadcast — every [`Receiver`] always sees
2//! the most recently sent value, never a backlog of every value sent
3//! (unlike [`super::broadcast`]).
4
5use super::wait_queue::WaitQueue;
6use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
7use std::sync::{Arc, RwLock};
8use std::task::{Context, Poll};
9
10#[repr(align(64))]
11struct Shared<T> {
12    value: RwLock<T>,
13    /// Bumped on every `send`; a `Receiver` compares this against the
14    /// version it last observed to know whether the value has changed
15    /// since its last `changed()`/construction.
16    version: AtomicU64,
17    sender_count: AtomicUsize,
18    receiver_count: AtomicUsize,
19    wait: WaitQueue,
20}
21
22/// Create a watch channel seeded with `init`.
23#[must_use]
24#[inline]
25pub fn channel<T>(init: T) -> (Sender<T>, Receiver<T>) {
26    let shared = Arc::new(Shared {
27        value: RwLock::new(init),
28        version: AtomicU64::new(0),
29        sender_count: AtomicUsize::new(1),
30        receiver_count: AtomicUsize::new(1),
31        wait: WaitQueue::new(),
32    });
33    let receiver = Receiver {
34        shared: shared.clone(),
35        seen_version: 0,
36    };
37    (Sender { shared }, receiver)
38}
39
40/// The sending half of a [`channel`]. Cheaply [`Clone`]-able.
41#[repr(align(64))]
42pub struct Sender<T> {
43    shared: Arc<Shared<T>>,
44}
45
46impl<T> Clone for Sender<T> {
47    #[inline(always)]
48    fn clone(&self) -> Self {
49        self.shared.sender_count.fetch_add(1, Ordering::AcqRel);
50        Self {
51            shared: self.shared.clone(),
52        }
53    }
54}
55
56impl<T> Drop for Sender<T> {
57    #[inline(always)]
58    fn drop(&mut self) {
59        if self.shared.sender_count.fetch_sub(1, Ordering::AcqRel) == 1 {
60            // Deliberately does NOT bump `version` — closing isn't a
61            // value change, and `poll_changed` checks `is_closed()` only
62            // *after* `try_observe_change()` finds nothing new, so
63            // bumping the version here would make a plain close look
64            // like an unseen value change and return `Ok(())` instead of
65            // `Err`. Just wake any waiter blocked in `changed()` so it
66            // re-polls, observes `sender_count == 0` via `is_closed()`,
67            // and returns the correct `Err` — rather than waiting forever
68            // for a send that will never come.
69            self.shared.wait.wake_all();
70        }
71    }
72}
73
74impl<T> Sender<T> {
75    /// Replace the current value with `value`, notifying every receiver.
76    #[inline(always)]
77    pub fn send(&self, value: T) {
78        *self
79            .shared
80            .value
81            .write()
82            .unwrap_or_else(std::sync::PoisonError::into_inner) = value;
83        self.shared.version.fetch_add(1, Ordering::Release);
84        self.shared.wait.wake_all();
85    }
86
87    /// A read-only snapshot of the current value.
88    #[inline(always)]
89    pub fn borrow(&self) -> std::sync::RwLockReadGuard<'_, T> {
90        self.shared
91            .value
92            .read()
93            .unwrap_or_else(std::sync::PoisonError::into_inner)
94    }
95
96    /// `true` once every [`Receiver`] has been dropped.
97    #[must_use]
98    #[inline(always)]
99    pub fn is_closed(&self) -> bool {
100        self.shared.receiver_count.load(Ordering::Acquire) == 0
101    }
102}
103
104/// The receiving half of a [`channel`].
105///
106/// [`Clone`]-able — each clone tracks its own "last seen version"
107/// independently, so every receiver (original and clones) sees every
108/// value change via its own [`changed`](Self::changed) calls.
109#[repr(align(64))]
110pub struct Receiver<T> {
111    shared: Arc<Shared<T>>,
112    seen_version: u64,
113}
114
115impl<T> Clone for Receiver<T> {
116    #[inline(always)]
117    fn clone(&self) -> Self {
118        self.shared.receiver_count.fetch_add(1, Ordering::AcqRel);
119        Self {
120            shared: self.shared.clone(),
121            seen_version: self.seen_version,
122        }
123    }
124}
125
126impl<T> Drop for Receiver<T> {
127    #[inline(always)]
128    fn drop(&mut self) {
129        self.shared.receiver_count.fetch_sub(1, Ordering::AcqRel);
130    }
131}
132
133impl<T: Send + Sync> Receiver<T> {
134    /// Wait until the value has changed since the last time this
135    /// receiver observed it (via construction or a previous
136    /// `changed()`/`borrow_and_update()`), then mark it seen.
137    ///
138    /// # Errors
139    /// Returns [`RecvError`] once every [`Sender`] has been dropped and
140    /// there are no further changes to observe.
141    #[inline(always)]
142    pub async fn changed(&mut self) -> Result<(), RecvError> {
143        std::future::poll_fn(|cx| self.poll_changed(cx)).await
144    }
145}
146
147impl<T> Receiver<T> {
148    /// A read-only snapshot of the current value. Does not mark it as
149    /// "seen" — a subsequent [`changed`](Self::changed) still resolves
150    /// immediately if the value changed before this call.
151    #[inline(always)]
152    pub fn borrow(&self) -> std::sync::RwLockReadGuard<'_, T> {
153        self.shared
154            .value
155            .read()
156            .unwrap_or_else(std::sync::PoisonError::into_inner)
157    }
158
159    /// Like [`borrow`](Self::borrow), but also marks the current value as
160    /// seen — a subsequent [`changed`](Self::changed) only resolves on a
161    /// value sent *after* this call.
162    #[inline(always)]
163    pub fn borrow_and_update(&mut self) -> std::sync::RwLockReadGuard<'_, T> {
164        self.seen_version = self.shared.version.load(Ordering::Acquire);
165        self.borrow()
166    }
167
168    #[inline]
169    fn poll_changed(&mut self, cx: &Context<'_>) -> Poll<Result<(), RecvError>> {
170        if self.try_observe_change() {
171            return Poll::Ready(Ok(()));
172        }
173        if self.is_closed() {
174            // The last sender could have sent a final value and then
175            // dropped concurrently with the check above — see
176            // `mpsc::Receiver::poll_recv`'s identical comment for why one
177            // more observe attempt is needed before reporting closed.
178            return Poll::Ready(if self.try_observe_change() {
179                Ok(())
180            } else {
181                Err(RecvError)
182            });
183        }
184        let token = self.shared.wait.register(cx.waker());
185        if self.try_observe_change() {
186            self.shared.wait.cancel(token);
187            return Poll::Ready(Ok(()));
188        }
189        if self.is_closed() {
190            let result = if self.try_observe_change() {
191                Ok(())
192            } else {
193                Err(RecvError)
194            };
195            self.shared.wait.cancel(token);
196            return Poll::Ready(result);
197        }
198        Poll::Pending
199    }
200
201    #[inline(always)]
202    fn try_observe_change(&mut self) -> bool {
203        let current = self.shared.version.load(Ordering::Acquire);
204        if current == self.seen_version {
205            false
206        } else {
207            self.seen_version = current;
208            true
209        }
210    }
211
212    #[inline(always)]
213    fn is_closed(&self) -> bool {
214        self.shared.sender_count.load(Ordering::Acquire) == 0
215    }
216}
217
218/// Error returned by [`Receiver::changed`] once every [`Sender`] has been
219/// dropped.
220#[derive(Debug, Clone, Copy, PartialEq, Eq)]
221#[repr(align(64))]
222pub struct RecvError;
223
224impl std::fmt::Display for RecvError {
225    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
226        f.write_str("channel closed: every sender dropped")
227    }
228}
229
230impl std::error::Error for RecvError {}