Skip to main content

poise_discovery/
snapshot.rs

1use std::{
2    error::Error,
3    fmt,
4    ops::Deref,
5    pin::Pin,
6    sync::{
7        Arc, Mutex, MutexGuard, Weak,
8        atomic::{AtomicBool, Ordering},
9    },
10    task::{Context, Poll, Waker},
11};
12
13use arc_swap::ArcSwap;
14use futures_core::Stream;
15
16use crate::Revision;
17
18/// An immutable, coherent view of discovered members.
19#[derive(Debug)]
20pub struct Snapshot<T> {
21    revision: Revision,
22    members: Arc<[T]>,
23}
24
25impl<T> Snapshot<T> {
26    /// Creates a snapshot from an owned member vector.
27    #[must_use]
28    pub fn new(revision: Revision, members: Vec<T>) -> Self {
29        Self {
30            revision,
31            members: members.into(),
32        }
33    }
34
35    /// Creates an empty initial snapshot.
36    #[must_use]
37    pub fn empty() -> Self {
38        Self::new(Revision::INITIAL, Vec::new())
39    }
40
41    /// Returns the membership revision.
42    #[must_use]
43    pub const fn revision(&self) -> Revision {
44        self.revision
45    }
46
47    /// Returns the snapshot members.
48    #[must_use]
49    pub fn members(&self) -> &[T] {
50        &self.members
51    }
52
53    /// Returns a shared handle to the member allocation.
54    #[must_use]
55    pub fn members_arc(&self) -> Arc<[T]> {
56        Arc::clone(&self.members)
57    }
58
59    /// Returns the number of members, including draining members.
60    #[must_use]
61    pub fn len(&self) -> usize {
62        self.members.len()
63    }
64
65    /// Returns whether the snapshot has no active or draining members.
66    #[must_use]
67    pub fn is_empty(&self) -> bool {
68        self.members.is_empty()
69    }
70}
71
72impl<T> Default for Snapshot<T> {
73    fn default() -> Self {
74        Self::empty()
75    }
76}
77
78impl<T> Deref for Snapshot<T> {
79    type Target = [T];
80
81    fn deref(&self) -> &Self::Target {
82        self.members()
83    }
84}
85
86impl<T> AsRef<[T]> for Snapshot<T> {
87    fn as_ref(&self) -> &[T] {
88        self.members()
89    }
90}
91
92/// Creates a single-writer publisher and cloneable atomic readers.
93#[must_use]
94pub fn snapshot_channel<T>(initial: Snapshot<T>) -> (SnapshotPublisher<T>, SnapshotReader<T>) {
95    let shared = Arc::new(Shared {
96        current: ArcSwap::from_pointee(initial),
97        publisher_alive: AtomicBool::new(true),
98        waiters: Mutex::new(Vec::new()),
99    });
100    (
101        SnapshotPublisher {
102            shared: Arc::clone(&shared),
103        },
104        SnapshotReader { shared },
105    )
106}
107
108struct Waiter {
109    waker: Mutex<Option<Waker>>,
110}
111
112struct Shared<T> {
113    current: ArcSwap<Snapshot<T>>,
114    publisher_alive: AtomicBool,
115    waiters: Mutex<Vec<Weak<Waiter>>>,
116}
117
118impl<T> Shared<T> {
119    fn wake_subscribers(&self) {
120        let mut wake = Vec::new();
121        self.lock_waiters().retain(|weak| {
122            let Some(waiter) = weak.upgrade() else {
123                return false;
124            };
125            if let Some(waker) = lock_waker(&waiter).take() {
126                wake.push(waker);
127            }
128            true
129        });
130        for waker in wake {
131            waker.wake();
132        }
133    }
134
135    fn lock_waiters(&self) -> MutexGuard<'_, Vec<Weak<Waiter>>> {
136        self.waiters
137            .lock()
138            .unwrap_or_else(std::sync::PoisonError::into_inner)
139    }
140}
141
142fn lock_waker(waiter: &Waiter) -> MutexGuard<'_, Option<Waker>> {
143    waiter
144        .waker
145        .lock()
146        .unwrap_or_else(std::sync::PoisonError::into_inner)
147}
148
149/// The single-writer side of an atomic snapshot cell.
150///
151/// Publishing requires `&mut self`, making the intended single-writer model
152/// explicit. This type is deliberately not `Clone`.
153pub struct SnapshotPublisher<T> {
154    shared: Arc<Shared<T>>,
155}
156
157impl<T> SnapshotPublisher<T> {
158    /// Atomically publishes a strictly newer snapshot.
159    ///
160    /// # Errors
161    ///
162    /// Returns [`PublishError`] when the attempted revision is not newer than
163    /// the currently published revision.
164    pub fn publish(&mut self, snapshot: Snapshot<T>) -> Result<Arc<Snapshot<T>>, PublishError> {
165        let current = self.shared.current.load_full();
166        if snapshot.revision() <= current.revision() {
167            return Err(PublishError {
168                current: current.revision(),
169                attempted: snapshot.revision(),
170            });
171        }
172
173        let snapshot = Arc::new(snapshot);
174        self.shared.current.store(Arc::clone(&snapshot));
175        self.shared.wake_subscribers();
176        Ok(snapshot)
177    }
178
179    /// Returns the currently published revision.
180    #[must_use]
181    pub fn revision(&self) -> Revision {
182        self.shared.current.load().revision()
183    }
184}
185
186impl<T> Drop for SnapshotPublisher<T> {
187    fn drop(&mut self) {
188        self.shared.publisher_alive.store(false, Ordering::Release);
189        self.shared.wake_subscribers();
190    }
191}
192
193/// A cloneable reader for atomically published snapshots.
194pub struct SnapshotReader<T> {
195    shared: Arc<Shared<T>>,
196}
197
198impl<T> SnapshotReader<T> {
199    /// Loads one coherent snapshot.
200    #[must_use]
201    pub fn load(&self) -> Arc<Snapshot<T>> {
202        self.shared.current.load_full()
203    }
204
205    /// Returns the currently published revision.
206    #[must_use]
207    pub fn revision(&self) -> Revision {
208        self.shared.current.load().revision()
209    }
210
211    /// Subscribes to the current snapshot and coalesced future publications.
212    ///
213    /// The first poll yields the current snapshot. If several revisions are
214    /// published before the next poll, the stream yields only the newest one.
215    /// It ends after the publisher is dropped and the latest snapshot has been
216    /// observed.
217    #[must_use]
218    pub fn subscribe(&self) -> SnapshotStream<T> {
219        SnapshotStream::new(self.clone(), None)
220    }
221
222    /// Subscribes only to publications newer than the currently visible
223    /// revision.
224    ///
225    /// Like [`subscribe`](Self::subscribe), intermediate revisions may be
226    /// coalesced.
227    #[must_use]
228    pub fn changes(&self) -> SnapshotStream<T> {
229        SnapshotStream::new(self.clone(), Some(self.revision()))
230    }
231}
232
233impl<T> Clone for SnapshotReader<T> {
234    fn clone(&self) -> Self {
235        Self {
236            shared: Arc::clone(&self.shared),
237        }
238    }
239}
240
241/// A runtime-neutral stream of coherent atomic snapshots.
242///
243/// Each subscriber owns an independent revision cursor and waker. Publications
244/// are state, not an event log: a slow subscriber observes the latest revision
245/// and may skip intermediate snapshots.
246pub struct SnapshotStream<T> {
247    reader: SnapshotReader<T>,
248    waiter: Arc<Waiter>,
249    last_seen: Option<Revision>,
250    terminated: bool,
251}
252
253impl<T> SnapshotStream<T> {
254    fn new(reader: SnapshotReader<T>, last_seen: Option<Revision>) -> Self {
255        let waiter = Arc::new(Waiter {
256            waker: Mutex::new(None),
257        });
258        reader.shared.lock_waiters().push(Arc::downgrade(&waiter));
259        Self {
260            reader,
261            waiter,
262            last_seen,
263            terminated: false,
264        }
265    }
266
267    /// Returns the latest revision yielded by this subscriber.
268    #[must_use]
269    pub const fn last_seen(&self) -> Option<Revision> {
270        self.last_seen
271    }
272
273    /// Returns whether publisher closure has terminated this stream.
274    #[must_use]
275    pub const fn is_terminated(&self) -> bool {
276        self.terminated
277    }
278
279    /// Polls for the next coherent snapshot without requiring an extension
280    /// trait import.
281    pub fn poll_snapshot(&mut self, context: &mut Context<'_>) -> Poll<Option<Arc<Snapshot<T>>>> {
282        Stream::poll_next(Pin::new(self), context)
283    }
284}
285
286impl<T> Unpin for SnapshotStream<T> {}
287
288impl<T> fmt::Debug for SnapshotStream<T> {
289    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
290        f.debug_struct("SnapshotStream")
291            .field("last_seen", &self.last_seen)
292            .field("terminated", &self.terminated)
293            .finish_non_exhaustive()
294    }
295}
296
297impl<T> Stream for SnapshotStream<T> {
298    type Item = Arc<Snapshot<T>>;
299
300    fn poll_next(mut self: Pin<&mut Self>, context: &mut Context<'_>) -> Poll<Option<Self::Item>> {
301        if self.terminated {
302            return Poll::Ready(None);
303        }
304
305        let current = self.reader.load();
306        if self
307            .last_seen
308            .is_none_or(|revision| current.revision() > revision)
309        {
310            self.last_seen = Some(current.revision());
311            return Poll::Ready(Some(current));
312        }
313        if !self.reader.shared.publisher_alive.load(Ordering::Acquire) {
314            self.terminated = true;
315            return Poll::Ready(None);
316        }
317
318        {
319            let mut registered = lock_waker(&self.waiter);
320            if registered
321                .as_ref()
322                .is_none_or(|waker| !waker.will_wake(context.waker()))
323            {
324                *registered = Some(context.waker().clone());
325            }
326        }
327
328        // Close the check/register race with publication and publisher drop.
329        let current = self.reader.load();
330        if self
331            .last_seen
332            .is_some_and(|revision| current.revision() > revision)
333        {
334            lock_waker(&self.waiter).take();
335            self.last_seen = Some(current.revision());
336            return Poll::Ready(Some(current));
337        }
338        if !self.reader.shared.publisher_alive.load(Ordering::Acquire) {
339            lock_waker(&self.waiter).take();
340            self.terminated = true;
341            return Poll::Ready(None);
342        }
343
344        Poll::Pending
345    }
346}
347
348impl<T> fmt::Debug for SnapshotReader<T> {
349    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
350        f.debug_struct("SnapshotReader")
351            .field("revision", &self.revision())
352            .finish_non_exhaustive()
353    }
354}
355
356/// A rejected non-monotonic publication.
357#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
358pub struct PublishError {
359    current: Revision,
360    attempted: Revision,
361}
362
363impl PublishError {
364    /// Returns the revision visible when publication was attempted.
365    #[must_use]
366    pub const fn current(self) -> Revision {
367        self.current
368    }
369
370    /// Returns the rejected revision.
371    #[must_use]
372    pub const fn attempted(self) -> Revision {
373        self.attempted
374    }
375}
376
377impl fmt::Display for PublishError {
378    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
379        write!(
380            f,
381            "snapshot revision {} is not newer than published revision {}",
382            self.attempted, self.current
383        )
384    }
385}
386
387impl Error for PublishError {}
388
389#[cfg(test)]
390mod tests {
391    use std::{
392        sync::{
393            Arc,
394            atomic::{AtomicUsize, Ordering as AtomicOrdering},
395        },
396        task::Wake,
397        thread,
398    };
399
400    use super::*;
401
402    #[derive(Default)]
403    struct WakeCount(AtomicUsize);
404
405    impl Wake for WakeCount {
406        fn wake(self: Arc<Self>) {
407            self.0.fetch_add(1, AtomicOrdering::Relaxed);
408        }
409
410        fn wake_by_ref(self: &Arc<Self>) {
411            self.0.fetch_add(1, AtomicOrdering::Relaxed);
412        }
413    }
414
415    fn poll_stream<T>(
416        stream: &mut SnapshotStream<T>,
417        waker: &Waker,
418    ) -> Poll<Option<Arc<Snapshot<T>>>> {
419        Stream::poll_next(Pin::new(stream), &mut Context::from_waker(waker))
420    }
421
422    #[test]
423    fn publication_requires_a_strictly_newer_revision() {
424        let (mut publisher, reader) =
425            snapshot_channel(Snapshot::new(Revision::new(5), vec!["current"]));
426
427        for attempted in [4, 5] {
428            let error = publisher
429                .publish(Snapshot::new(Revision::new(attempted), vec!["stale"]))
430                .unwrap_err();
431            assert_eq!(error.current(), Revision::new(5));
432            assert_eq!(error.attempted(), Revision::new(attempted));
433        }
434        assert_eq!(reader.load().members(), ["current"]);
435
436        publisher
437            .publish(Snapshot::new(Revision::new(6), vec!["new"]))
438            .unwrap();
439        assert_eq!(reader.load().members(), ["new"]);
440    }
441
442    #[test]
443    fn concurrent_readers_never_observe_a_torn_snapshot() {
444        const MEMBER_COUNT: usize = 64;
445        const LAST_REVISION: u64 = 2_000;
446
447        let (mut publisher, reader) =
448            snapshot_channel(Snapshot::new(Revision::INITIAL, vec![0_u64; MEMBER_COUNT]));
449        let writer = thread::spawn(move || {
450            for revision in 1..=LAST_REVISION {
451                publisher
452                    .publish(Snapshot::new(
453                        Revision::new(revision),
454                        vec![revision; MEMBER_COUNT],
455                    ))
456                    .unwrap();
457            }
458        });
459
460        let readers: Vec<_> = (0..4)
461            .map(|_| {
462                let reader = reader.clone();
463                thread::spawn(move || {
464                    for _ in 0..10_000 {
465                        let snapshot = reader.load();
466                        let expected = snapshot.revision().get();
467                        assert_eq!(snapshot.len(), MEMBER_COUNT);
468                        assert!(snapshot.iter().all(|member| *member == expected));
469                    }
470                })
471            })
472            .collect();
473
474        writer.join().unwrap();
475        for reader in readers {
476            reader.join().unwrap();
477        }
478    }
479
480    #[test]
481    fn old_snapshot_handles_remain_valid_after_publication() {
482        let (mut publisher, reader) =
483            snapshot_channel(Snapshot::new(Revision::INITIAL, vec![1, 2]));
484        let old = reader.load();
485
486        publisher
487            .publish(Snapshot::new(Revision::new(1), vec![3]))
488            .unwrap();
489
490        assert_eq!(old.members(), [1, 2]);
491        assert_eq!(reader.load().members(), [3]);
492    }
493
494    #[test]
495    fn subscription_yields_current_then_wakes_for_a_newer_snapshot() {
496        let (mut publisher, reader) =
497            snapshot_channel(Snapshot::new(Revision::INITIAL, vec!["initial"]));
498        let mut stream = reader.subscribe();
499        let wake_count = Arc::new(WakeCount::default());
500        let waker = Waker::from(Arc::clone(&wake_count));
501
502        let Poll::Ready(Some(initial)) = poll_stream(&mut stream, &waker) else {
503            panic!("subscription did not yield its initial snapshot");
504        };
505        assert_eq!(initial.members(), ["initial"]);
506        assert_eq!(stream.last_seen(), Some(Revision::INITIAL));
507        assert!(poll_stream(&mut stream, &waker).is_pending());
508
509        publisher
510            .publish(Snapshot::new(Revision::new(1), vec!["new"]))
511            .unwrap();
512        assert_eq!(wake_count.0.load(AtomicOrdering::Relaxed), 1);
513        let Poll::Ready(Some(new)) = poll_stream(&mut stream, &waker) else {
514            panic!("subscription did not yield the published snapshot");
515        };
516        assert_eq!(new.revision(), Revision::new(1));
517        assert_eq!(new.members(), ["new"]);
518    }
519
520    #[test]
521    fn changes_waits_for_a_revision_newer_than_construction() {
522        let (mut publisher, reader) = snapshot_channel(Snapshot::new(Revision::new(4), vec![4]));
523        let mut changes = reader.changes();
524        let waker = Waker::noop();
525
526        assert_eq!(changes.last_seen(), Some(Revision::new(4)));
527        assert!(poll_stream(&mut changes, waker).is_pending());
528        publisher
529            .publish(Snapshot::new(Revision::new(5), vec![5]))
530            .unwrap();
531        let Poll::Ready(Some(snapshot)) = poll_stream(&mut changes, waker) else {
532            panic!("change stream did not yield the newer revision");
533        };
534        assert_eq!(snapshot.members(), [5]);
535    }
536
537    #[test]
538    fn slow_subscribers_coalesce_to_the_latest_revision() {
539        let (mut publisher, reader) = snapshot_channel(Snapshot::empty());
540        let mut changes = reader.changes();
541        publisher
542            .publish(Snapshot::new(Revision::new(1), vec![1]))
543            .unwrap();
544        publisher
545            .publish(Snapshot::new(Revision::new(2), vec![2]))
546            .unwrap();
547
548        let Poll::Ready(Some(snapshot)) = poll_stream(&mut changes, Waker::noop()) else {
549            panic!("change stream did not yield the latest snapshot");
550        };
551        assert_eq!(snapshot.revision(), Revision::new(2));
552        assert_eq!(snapshot.members(), [2]);
553        assert!(poll_stream(&mut changes, Waker::noop()).is_pending());
554    }
555
556    #[test]
557    fn publication_wakes_each_independent_subscriber() {
558        let (mut publisher, reader) = snapshot_channel(Snapshot::<usize>::empty());
559        let mut left = reader.changes();
560        let mut right = reader.changes();
561        let left_count = Arc::new(WakeCount::default());
562        let right_count = Arc::new(WakeCount::default());
563        let left_waker = Waker::from(Arc::clone(&left_count));
564        let right_waker = Waker::from(Arc::clone(&right_count));
565        assert!(poll_stream(&mut left, &left_waker).is_pending());
566        assert!(poll_stream(&mut right, &right_waker).is_pending());
567
568        publisher
569            .publish(Snapshot::new(Revision::new(1), vec![1]))
570            .unwrap();
571        assert_eq!(left_count.0.load(AtomicOrdering::Relaxed), 1);
572        assert_eq!(right_count.0.load(AtomicOrdering::Relaxed), 1);
573        assert!(matches!(
574            poll_stream(&mut left, &left_waker),
575            Poll::Ready(Some(_))
576        ));
577        assert!(matches!(
578            poll_stream(&mut right, &right_waker),
579            Poll::Ready(Some(_))
580        ));
581    }
582
583    #[test]
584    fn publisher_drop_wakes_and_terminates_subscribers_after_latest_state() {
585        let (publisher, reader) = snapshot_channel(Snapshot::new(Revision::INITIAL, vec!["last"]));
586        let mut stream = reader.subscribe();
587        let wake_count = Arc::new(WakeCount::default());
588        let waker = Waker::from(Arc::clone(&wake_count));
589        assert!(matches!(
590            poll_stream(&mut stream, &waker),
591            Poll::Ready(Some(_))
592        ));
593        assert!(poll_stream(&mut stream, &waker).is_pending());
594
595        drop(publisher);
596        assert_eq!(wake_count.0.load(AtomicOrdering::Relaxed), 1);
597        assert!(matches!(
598            poll_stream(&mut stream, &waker),
599            Poll::Ready(None)
600        ));
601        assert!(stream.is_terminated());
602        assert!(matches!(
603            poll_stream(&mut stream, &waker),
604            Poll::Ready(None)
605        ));
606    }
607
608    #[test]
609    fn rejected_publication_does_not_wake_subscribers() {
610        let (mut publisher, reader) = snapshot_channel(Snapshot::<usize>::empty());
611        let mut changes = reader.changes();
612        let wake_count = Arc::new(WakeCount::default());
613        let waker = Waker::from(Arc::clone(&wake_count));
614        assert!(poll_stream(&mut changes, &waker).is_pending());
615
616        assert!(
617            publisher
618                .publish(Snapshot::new(Revision::INITIAL, vec![1]))
619                .is_err()
620        );
621        assert_eq!(wake_count.0.load(AtomicOrdering::Relaxed), 0);
622        assert!(poll_stream(&mut changes, &waker).is_pending());
623    }
624}