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#[derive(Debug)]
20pub struct Snapshot<T> {
21 revision: Revision,
22 members: Arc<[T]>,
23}
24
25impl<T> Snapshot<T> {
26 #[must_use]
28 pub fn new(revision: Revision, members: Vec<T>) -> Self {
29 Self {
30 revision,
31 members: members.into(),
32 }
33 }
34
35 #[must_use]
37 pub fn empty() -> Self {
38 Self::new(Revision::INITIAL, Vec::new())
39 }
40
41 #[must_use]
43 pub const fn revision(&self) -> Revision {
44 self.revision
45 }
46
47 #[must_use]
49 pub fn members(&self) -> &[T] {
50 &self.members
51 }
52
53 #[must_use]
55 pub fn members_arc(&self) -> Arc<[T]> {
56 Arc::clone(&self.members)
57 }
58
59 #[must_use]
61 pub fn len(&self) -> usize {
62 self.members.len()
63 }
64
65 #[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#[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
149pub struct SnapshotPublisher<T> {
154 shared: Arc<Shared<T>>,
155}
156
157impl<T> SnapshotPublisher<T> {
158 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 #[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
193pub struct SnapshotReader<T> {
195 shared: Arc<Shared<T>>,
196}
197
198impl<T> SnapshotReader<T> {
199 #[must_use]
201 pub fn load(&self) -> Arc<Snapshot<T>> {
202 self.shared.current.load_full()
203 }
204
205 #[must_use]
207 pub fn revision(&self) -> Revision {
208 self.shared.current.load().revision()
209 }
210
211 #[must_use]
218 pub fn subscribe(&self) -> SnapshotStream<T> {
219 SnapshotStream::new(self.clone(), None)
220 }
221
222 #[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
241pub 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 #[must_use]
269 pub const fn last_seen(&self) -> Option<Revision> {
270 self.last_seen
271 }
272
273 #[must_use]
275 pub const fn is_terminated(&self) -> bool {
276 self.terminated
277 }
278
279 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 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#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq)]
358pub struct PublishError {
359 current: Revision,
360 attempted: Revision,
361}
362
363impl PublishError {
364 #[must_use]
366 pub const fn current(self) -> Revision {
367 self.current
368 }
369
370 #[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}