Skip to main content

doom_fish_utils/
stream.rs

1//! Executor-agnostic bounded async streams for FFI callbacks.
2//!
3//! `BoundedAsyncStream<T>` is a generic, runtime-agnostic stream primitive
4//! designed for wrapping Apple SDK callback / delegate / KVO patterns:
5//!
6//! * **Bounded** — backed by a fixed-capacity `VecDeque`. When the buffer
7//!   is full and a new item arrives from the producer, the **oldest**
8//!   queued item is dropped to make room (lossy by design).
9//! * **Waker-driven** — implements `std::future::Future` via a stored
10//!   `Waker`; works with any executor (tokio, async-std, smol, futures,
11//!   etc.) without requiring a runtime feature.
12//! * **`Send + Sync`** — produces and consumes can live on different
13//!   threads, locked by a single `Mutex`.
14//!
15//! The lossy-oldest-drop policy is the right default for real-time event
16//! streams (UI input, frame capture, BLE notifications, location updates):
17//! a slow consumer should always see the latest event, not a stale queue.
18//! When you instead need back-pressure (every event must be delivered),
19//! use [`AsyncStreamSender::push_or_block`] which blocks the producer
20//! until the consumer drains capacity.
21//!
22//! # Example
23//!
24//! ```no_run
25//! use doom_fish_utils::stream::BoundedAsyncStream;
26//! use std::sync::Arc;
27//!
28//! # async fn run() {
29//! // 8-element ring buffer of `String` events.
30//! let (stream, sender) = BoundedAsyncStream::<String>::new(8);
31//!
32//! // Producer side: typically a Swift delegate / extern "C" callback
33//! // running on a background queue.
34//! std::thread::spawn(move || {
35//!     for i in 0..100 {
36//!         sender.push(format!("event #{i}"));
37//!     }
38//!     drop(sender); // closes the stream
39//! });
40//!
41//! // Consumer side: any async runtime.
42//! while let Some(event) = stream.next().await {
43//!     println!("got {event}");
44//! }
45//! # }
46//! ```
47
48use std::collections::VecDeque;
49use std::fmt;
50use std::future::Future;
51use std::pin::Pin;
52use std::sync::{Arc, Condvar, Mutex, MutexGuard};
53use std::task::{Context, Poll, Waker};
54
55/// Backing storage shared between the [`BoundedAsyncStream`] consumer and
56/// every [`AsyncStreamSender`] producer.
57struct State<T> {
58    buffer: VecDeque<T>,
59    waiters: Vec<(u64, Waker)>,
60    next_waiter: u64,
61    /// Set to `true` when every sender has been dropped. The consumer's
62    /// `next()` then returns `None` once the buffer drains.
63    closed: bool,
64    /// Set to `true` when the stream is dropped — wakes any blocked
65    /// producers so they can bail out instead of waiting forever.
66    consumer_gone: bool,
67    sender_count: usize,
68    #[cfg(test)]
69    blocked_producers: usize,
70}
71
72#[cfg(feature = "futures-stream")]
73const STREAM_WAITER: u64 = 0;
74
75impl<T> State<T> {
76    fn register_waiter(
77        &mut self,
78        waiter: &mut Option<u64>,
79        current: &Waker,
80        next_waker: &mut Option<Waker>,
81    ) -> Option<Waker> {
82        let id = *waiter.get_or_insert_with(|| {
83            let id = self.next_waiter;
84            self.next_waiter = id.wrapping_add(1).max(1);
85            id
86        });
87        match self
88            .waiters
89            .iter_mut()
90            .find(|(registered, _)| *registered == id)
91        {
92            Some((_, existing)) if existing.will_wake(current) => None,
93            Some((_, existing)) => next_waker
94                .take()
95                .map(|waker| std::mem::replace(existing, waker)),
96            None => {
97                self.waiters
98                    .extend(next_waker.take().map(|waker| (id, waker)));
99                None
100            }
101        }
102    }
103
104    fn remove_waiter(&mut self, waiter: Option<u64>) -> Option<Waker> {
105        let waiter = waiter?;
106        let index = self
107            .waiters
108            .iter()
109            .position(|(registered, _)| *registered == waiter)?;
110        Some(self.waiters.swap_remove(index).1)
111    }
112}
113
114fn wake_waiters(waiters: Vec<(u64, Waker)>) {
115    for (_, waker) in waiters {
116        waker.wake();
117    }
118}
119
120struct Shared<T> {
121    state: Mutex<State<T>>,
122    capacity_available: Condvar,
123    capacity: usize,
124}
125
126impl<T> Shared<T> {
127    fn lock_state(&self) -> MutexGuard<'_, State<T>> {
128        self.state
129            .lock()
130            .unwrap_or_else(|_| panic!("BoundedAsyncStream state mutex poisoned"))
131    }
132
133    fn lock_state_for_drop(&self) -> MutexGuard<'_, State<T>> {
134        self.state
135            .lock()
136            .unwrap_or_else(std::sync::PoisonError::into_inner)
137    }
138
139    #[cfg(test)]
140    fn wait_for_blocked_producers(&self, expected: usize) {
141        let (state, timeout) = self
142            .capacity_available
143            .wait_timeout_while(
144                self.lock_state(),
145                std::time::Duration::from_secs(5),
146                |state| state.blocked_producers < expected,
147            )
148            .unwrap_or_else(|_| panic!("BoundedAsyncStream state mutex poisoned"));
149        assert!(
150            !timeout.timed_out(),
151            "timed out waiting for {expected} blocked producer(s)"
152        );
153        drop(state);
154    }
155}
156
157/// A bounded, lossy-by-default, executor-agnostic async stream.
158///
159/// Items are pushed by one or more [`AsyncStreamSender`] handles and pulled
160/// asynchronously via [`BoundedAsyncStream::next`].
161///
162/// See the [module-level docs](crate::stream) for the full design rationale.
163pub struct BoundedAsyncStream<T> {
164    shared: Arc<Shared<T>>,
165}
166
167/// Producer handle for a [`BoundedAsyncStream`].
168///
169/// Cheap to clone (`Arc` under the hood). Drop the last `AsyncStreamSender`
170/// to close the stream; the consumer's `next()` will yield `None` once the
171/// buffer is empty.
172pub struct AsyncStreamSender<T> {
173    shared: Arc<Shared<T>>,
174}
175
176impl<T> Clone for AsyncStreamSender<T> {
177    fn clone(&self) -> Self {
178        let mut state = self.shared.lock_state();
179        state.sender_count = state
180            .sender_count
181            .checked_add(1)
182            .expect("AsyncStreamSender count overflow");
183        drop(state);
184
185        Self {
186            shared: Arc::clone(&self.shared),
187        }
188    }
189}
190
191impl<T> fmt::Debug for BoundedAsyncStream<T> {
192    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
193        f.debug_struct("BoundedAsyncStream")
194            .field("buffered", &self.buffered_count())
195            .field("capacity", &self.capacity())
196            .field("is_closed", &self.is_closed())
197            .finish_non_exhaustive()
198    }
199}
200
201impl<T> fmt::Debug for AsyncStreamSender<T> {
202    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
203        f.debug_struct("AsyncStreamSender").finish_non_exhaustive()
204    }
205}
206
207impl<T> BoundedAsyncStream<T> {
208    /// Creates a new bounded stream with the given capacity.
209    ///
210    /// Returns the consumer side and a single producer; clone the sender
211    /// to fan out to multiple producers.
212    ///
213    /// # Panics
214    ///
215    /// Panics if `capacity` is 0 — a zero-capacity buffer would drop every
216    /// item before the consumer could observe it. Use capacity 1 if you
217    /// genuinely want "latest only" semantics.
218    #[must_use]
219    pub fn new(capacity: usize) -> (Self, AsyncStreamSender<T>) {
220        assert!(capacity > 0, "BoundedAsyncStream capacity must be > 0");
221
222        let shared = Arc::new(Shared {
223            capacity,
224            capacity_available: Condvar::new(),
225            state: Mutex::new(State {
226                buffer: VecDeque::with_capacity(capacity),
227                waiters: Vec::new(),
228                next_waiter: 1,
229                closed: false,
230                consumer_gone: false,
231                sender_count: 1,
232                #[cfg(test)]
233                blocked_producers: 0,
234            }),
235        });
236
237        let stream = Self {
238            shared: Arc::clone(&shared),
239        };
240        let sender = AsyncStreamSender { shared };
241        (stream, sender)
242    }
243
244    /// Returns a future that resolves to the next item, or `None` once the
245    /// stream is closed and drained.
246    #[must_use]
247    pub const fn next(&self) -> NextItem<'_, T> {
248        NextItem {
249            stream: self,
250            waiter: None,
251        }
252    }
253
254    /// Non-blocking pop. Returns `None` if the buffer is empty (regardless
255    /// of whether the stream is open or closed).
256    ///
257    /// # Panics
258    ///
259    /// Panics if the shared state mutex is poisoned.
260    #[must_use]
261    pub fn try_next(&self) -> Option<T> {
262        let item = {
263            let mut state = self.shared.lock_state();
264            state.buffer.pop_front()
265        };
266        if item.is_some() {
267            self.shared.capacity_available.notify_one();
268        }
269        item
270    }
271
272    /// Returns `true` if the stream has been closed (all senders dropped).
273    /// Note: a closed stream may still have buffered items to drain.
274    ///
275    /// # Panics
276    ///
277    /// Panics if the shared state mutex is poisoned.
278    #[must_use]
279    pub fn is_closed(&self) -> bool {
280        self.shared.lock_state().closed
281    }
282
283    /// Returns the number of items currently buffered (0..=capacity).
284    ///
285    /// # Panics
286    ///
287    /// Panics if the shared state mutex is poisoned.
288    #[must_use]
289    pub fn buffered_count(&self) -> usize {
290        self.shared.lock_state().buffer.len()
291    }
292
293    /// Returns the buffer capacity, as passed to [`Self::new`].
294    #[must_use]
295    pub fn capacity(&self) -> usize {
296        self.shared.capacity
297    }
298
299    /// Drops all currently buffered items without closing the stream.
300    ///
301    /// # Panics
302    ///
303    /// Panics if the shared state mutex is poisoned.
304    pub fn clear_buffer(&self) {
305        let mut cleared = VecDeque::with_capacity(self.shared.capacity);
306        {
307            let mut state = self.shared.lock_state();
308            std::mem::swap(&mut state.buffer, &mut cleared);
309        }
310        self.shared.capacity_available.notify_all();
311        drop(cleared);
312    }
313
314    fn poll_next_item(&self, cx: &Context<'_>, waiter: &mut Option<u64>) -> Poll<Option<T>> {
315        let mut next_waker = Some(cx.waker().clone());
316        let (result, released_waker, made_room) = {
317            let mut state = self.shared.lock_state();
318
319            let outcome = match state.buffer.pop_front() {
320                Some(item) => (
321                    Poll::Ready(Some(item)),
322                    state.remove_waiter(waiter.take()),
323                    true,
324                ),
325                None if state.closed => {
326                    (Poll::Ready(None), state.remove_waiter(waiter.take()), false)
327                }
328                None => {
329                    let released = state.register_waiter(waiter, cx.waker(), &mut next_waker);
330                    (Poll::Pending, released, false)
331                }
332            };
333            drop(state);
334            outcome
335        };
336
337        drop(next_waker);
338        drop(released_waker);
339        if made_room {
340            self.shared.capacity_available.notify_one();
341        }
342        result
343    }
344}
345
346impl<T> Drop for BoundedAsyncStream<T> {
347    fn drop(&mut self) {
348        let stale_wakers = {
349            let mut state = self.shared.lock_state_for_drop();
350            state.consumer_gone = true;
351            std::mem::take(&mut state.waiters)
352        };
353        self.shared.capacity_available.notify_all();
354        drop(stale_wakers);
355    }
356}
357
358impl<T> AsyncStreamSender<T> {
359    /// Push an item; drops the oldest queued item if the buffer is at
360    /// capacity. This is the lossy default.
361    ///
362    /// # Panics
363    ///
364    /// Panics if the shared state mutex is poisoned.
365    pub fn push(&self, item: T) {
366        let (overwritten, waiters) = {
367            let mut state = self.shared.lock_state();
368            let overwritten = if state.buffer.len() >= self.shared.capacity {
369                state.buffer.pop_front()
370            } else {
371                None
372            };
373            state.buffer.push_back(item);
374            (overwritten, std::mem::take(&mut state.waiters))
375        };
376
377        wake_waiters(waiters);
378        drop(overwritten);
379    }
380
381    /// Push an item, blocking the current thread if the buffer is full
382    /// until the consumer drains an item.
383    ///
384    /// Returns `Err(item)` if the consumer side has been dropped — the
385    /// item is returned to the caller so it isn't leaked.
386    ///
387    /// # Errors
388    ///
389    /// Returns `Err(item)` if the consumer has been dropped.
390    ///
391    /// # Panics
392    ///
393    /// Panics if the shared state mutex is poisoned.
394    pub fn push_or_block(&self, item: T) -> Result<(), T> {
395        let mut state = self.shared.lock_state();
396        if !state.consumer_gone && state.buffer.len() >= self.shared.capacity {
397            #[cfg(test)]
398            {
399                state.blocked_producers += 1;
400                self.shared.capacity_available.notify_all();
401            }
402
403            state = self
404                .shared
405                .capacity_available
406                .wait_while(state, |state| {
407                    !state.consumer_gone && state.buffer.len() >= self.shared.capacity
408                })
409                .unwrap_or_else(|_| panic!("BoundedAsyncStream state mutex poisoned"));
410
411            #[cfg(test)]
412            {
413                state.blocked_producers -= 1;
414                self.shared.capacity_available.notify_all();
415            }
416        }
417
418        if state.consumer_gone {
419            drop(state);
420            return Err(item);
421        }
422
423        state.buffer.push_back(item);
424        let waiters = std::mem::take(&mut state.waiters);
425        drop(state);
426
427        wake_waiters(waiters);
428        Ok(())
429    }
430
431    /// Returns the number of items currently buffered.
432    ///
433    /// # Panics
434    ///
435    /// Panics if the shared state mutex is poisoned.
436    #[must_use]
437    pub fn buffered_count(&self) -> usize {
438        self.shared.lock_state().buffer.len()
439    }
440
441    /// Returns `true` if the consumer has been dropped.
442    ///
443    /// # Panics
444    ///
445    /// Panics if the shared state mutex is poisoned.
446    #[must_use]
447    pub fn is_consumer_gone(&self) -> bool {
448        self.shared.lock_state().consumer_gone
449    }
450}
451
452impl<T> Drop for AsyncStreamSender<T> {
453    fn drop(&mut self) {
454        let waiters = {
455            let mut state = self.shared.lock_state_for_drop();
456            state.sender_count -= 1;
457            if state.sender_count == 0 {
458                state.closed = true;
459                std::mem::take(&mut state.waiters)
460            } else {
461                Vec::new()
462            }
463        };
464
465        wake_waiters(waiters);
466    }
467}
468
469/// Future returned by [`BoundedAsyncStream::next`].
470pub struct NextItem<'a, T> {
471    stream: &'a BoundedAsyncStream<T>,
472    waiter: Option<u64>,
473}
474
475impl<T> fmt::Debug for NextItem<'_, T> {
476    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
477        f.debug_struct("NextItem").finish_non_exhaustive()
478    }
479}
480
481impl<T> Future for NextItem<'_, T> {
482    type Output = Option<T>;
483
484    fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Self::Output> {
485        let this = self.get_mut();
486        this.stream.poll_next_item(cx, &mut this.waiter)
487    }
488}
489
490impl<T> Drop for NextItem<'_, T> {
491    fn drop(&mut self) {
492        if self.waiter.is_some() {
493            let released = {
494                let mut state = self.stream.shared.lock_state_for_drop();
495                state.remove_waiter(self.waiter.take())
496            };
497            drop(released);
498        }
499    }
500}
501
502#[cfg(feature = "futures-stream")]
503impl<T: 'static> futures_core::Stream for BoundedAsyncStream<T> {
504    type Item = T;
505
506    fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<T>> {
507        self.poll_next_item(cx, &mut Some(STREAM_WAITER))
508    }
509}
510
511#[cfg(test)]
512mod tests {
513    #[cfg(feature = "futures-stream")]
514    use std::future::poll_fn;
515    use std::future::Future;
516    use std::panic::{catch_unwind, AssertUnwindSafe};
517    use std::pin::Pin;
518    use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
519    use std::sync::{mpsc, Arc, Barrier, TryLockError, Weak};
520    use std::task::{Context, Poll, Wake, Waker};
521    use std::thread::{self, JoinHandle};
522    use std::time::Duration;
523
524    #[cfg(feature = "futures-stream")]
525    use futures_core::Stream;
526
527    use super::{AsyncStreamSender, BoundedAsyncStream, Shared};
528
529    const TEST_TIMEOUT: Duration = Duration::from_secs(5);
530
531    fn spawn_blocking_push(
532        sender: AsyncStreamSender<u32>,
533        item: u32,
534    ) -> (mpsc::Receiver<Result<(), u32>>, JoinHandle<()>) {
535        let (result_tx, result_rx) = mpsc::channel();
536        let handle = thread::spawn(move || {
537            let result = sender.push_or_block(item);
538            result_tx.send(result).unwrap();
539        });
540        (result_rx, handle)
541    }
542
543    #[test]
544    fn try_next_notifies_blocked_producer() {
545        let (stream, sender) = BoundedAsyncStream::new(1);
546        sender.push(1);
547        let (result_rx, handle) = spawn_blocking_push(sender, 2);
548        stream.shared.wait_for_blocked_producers(1);
549
550        assert_eq!(stream.try_next(), Some(1));
551        assert_eq!(result_rx.recv_timeout(TEST_TIMEOUT).unwrap(), Ok(()));
552        handle.join().unwrap();
553        assert_eq!(stream.try_next(), Some(2));
554    }
555
556    #[test]
557    fn next_future_notifies_blocked_producer() {
558        let (stream, sender) = BoundedAsyncStream::new(1);
559        sender.push(1);
560        let (result_rx, handle) = spawn_blocking_push(sender, 2);
561        stream.shared.wait_for_blocked_producers(1);
562
563        assert_eq!(pollster::block_on(stream.next()), Some(1));
564        assert_eq!(result_rx.recv_timeout(TEST_TIMEOUT).unwrap(), Ok(()));
565        handle.join().unwrap();
566        assert_eq!(pollster::block_on(stream.next()), Some(2));
567    }
568
569    #[cfg(feature = "futures-stream")]
570    #[test]
571    fn stream_poll_notifies_blocked_producer() {
572        let (mut stream, sender) = BoundedAsyncStream::new(1);
573        sender.push(1);
574        let (result_rx, handle) = spawn_blocking_push(sender, 2);
575        stream.shared.wait_for_blocked_producers(1);
576
577        let first = pollster::block_on(poll_fn(|cx| Pin::new(&mut stream).poll_next(cx)));
578        assert_eq!(first, Some(1));
579        assert_eq!(result_rx.recv_timeout(TEST_TIMEOUT).unwrap(), Ok(()));
580        handle.join().unwrap();
581
582        let second = pollster::block_on(poll_fn(|cx| Pin::new(&mut stream).poll_next(cx)));
583        assert_eq!(second, Some(2));
584    }
585
586    #[test]
587    fn clear_buffer_notifies_blocked_producer() {
588        let (stream, sender) = BoundedAsyncStream::new(1);
589        sender.push(1);
590        let (result_rx, handle) = spawn_blocking_push(sender, 2);
591        stream.shared.wait_for_blocked_producers(1);
592
593        stream.clear_buffer();
594        assert_eq!(result_rx.recv_timeout(TEST_TIMEOUT).unwrap(), Ok(()));
595        handle.join().unwrap();
596        assert_eq!(stream.try_next(), Some(2));
597    }
598
599    #[test]
600    fn consumer_drop_returns_blocked_item() {
601        let (stream, sender) = BoundedAsyncStream::new(1);
602        sender.push(1);
603        let (result_rx, handle) = spawn_blocking_push(sender, 2);
604        stream.shared.wait_for_blocked_producers(1);
605
606        drop(stream);
607        assert_eq!(result_rx.recv_timeout(TEST_TIMEOUT).unwrap(), Err(2));
608        handle.join().unwrap();
609    }
610
611    #[test]
612    fn concurrent_sender_clone_drops_close_stream() {
613        let (stream, sender) = BoundedAsyncStream::<u32>::new(1);
614        let sender_clone = sender.clone();
615        let barrier = Arc::new(Barrier::new(3));
616
617        let first_barrier = Arc::clone(&barrier);
618        let first = thread::spawn(move || {
619            first_barrier.wait();
620            drop(sender);
621        });
622        let second_barrier = Arc::clone(&barrier);
623        let second = thread::spawn(move || {
624            second_barrier.wait();
625            drop(sender_clone);
626        });
627
628        barrier.wait();
629        first.join().unwrap();
630        second.join().unwrap();
631
632        assert!(stream.is_closed());
633        assert_eq!(pollster::block_on(stream.next()), None);
634    }
635
636    fn shared_state_is_unlocked<T>(shared: &Weak<Shared<T>>) -> bool {
637        let Some(shared) = shared.upgrade() else {
638            return true;
639        };
640        let unlocked = match shared.state.try_lock() {
641            Ok(_state) => true,
642            Err(TryLockError::WouldBlock | TryLockError::Poisoned(_)) => false,
643        };
644        unlocked
645    }
646
647    struct ReentrantWaker {
648        shared: Weak<Shared<u32>>,
649        wake_was_unlocked: Arc<AtomicBool>,
650        drop_was_unlocked: Arc<AtomicBool>,
651    }
652
653    impl Wake for ReentrantWaker {
654        fn wake(self: Arc<Self>) {
655            self.wake_was_unlocked
656                .store(shared_state_is_unlocked(&self.shared), Ordering::SeqCst);
657        }
658    }
659
660    impl Drop for ReentrantWaker {
661        fn drop(&mut self) {
662            self.drop_was_unlocked
663                .store(shared_state_is_unlocked(&self.shared), Ordering::SeqCst);
664        }
665    }
666
667    #[test]
668    fn wakes_and_drops_waker_outside_state_mutex() {
669        let (stream, sender) = BoundedAsyncStream::new(1);
670        let wake_was_unlocked = Arc::new(AtomicBool::new(false));
671        let drop_was_unlocked = Arc::new(AtomicBool::new(false));
672        let probe = Arc::new(ReentrantWaker {
673            shared: Arc::downgrade(&stream.shared),
674            wake_was_unlocked: Arc::clone(&wake_was_unlocked),
675            drop_was_unlocked: Arc::clone(&drop_was_unlocked),
676        });
677        let waker = Waker::from(Arc::clone(&probe));
678        let mut next = stream.next();
679
680        {
681            let mut cx = Context::from_waker(&waker);
682            assert_eq!(Pin::new(&mut next).poll(&mut cx), Poll::Pending);
683        }
684        drop(waker);
685        drop(probe);
686
687        sender.push(1);
688
689        assert!(wake_was_unlocked.load(Ordering::SeqCst));
690        assert!(drop_was_unlocked.load(Ordering::SeqCst));
691    }
692
693    struct ReentrantItem {
694        shared: Weak<Shared<Self>>,
695        unlocked_drops: Arc<AtomicUsize>,
696        locked_drops: Arc<AtomicUsize>,
697    }
698
699    impl Drop for ReentrantItem {
700        fn drop(&mut self) {
701            if shared_state_is_unlocked(&self.shared) {
702                self.unlocked_drops.fetch_add(1, Ordering::SeqCst);
703            } else {
704                self.locked_drops.fetch_add(1, Ordering::SeqCst);
705            }
706        }
707    }
708
709    #[test]
710    fn overwritten_and_cleared_items_drop_outside_state_mutex() {
711        let (stream, sender) = BoundedAsyncStream::new(1);
712        let unlocked_drops = Arc::new(AtomicUsize::new(0));
713        let locked_drops = Arc::new(AtomicUsize::new(0));
714
715        let make_item = || ReentrantItem {
716            shared: Arc::downgrade(&stream.shared),
717            unlocked_drops: Arc::clone(&unlocked_drops),
718            locked_drops: Arc::clone(&locked_drops),
719        };
720
721        sender.push(make_item());
722        sender.push(make_item());
723        assert_eq!(unlocked_drops.load(Ordering::SeqCst), 1);
724
725        stream.clear_buffer();
726        assert_eq!(unlocked_drops.load(Ordering::SeqCst), 2);
727        assert_eq!(locked_drops.load(Ordering::SeqCst), 0);
728    }
729
730    #[test]
731    fn poisoned_state_does_not_masquerade_as_close_or_delivery() {
732        let (stream, sender) = BoundedAsyncStream::new(1);
733        let shared = Arc::clone(&stream.shared);
734        assert!(thread::spawn(move || {
735            let _state = shared.state.lock().unwrap();
736            panic!("poison stream state");
737        })
738        .join()
739        .is_err());
740
741        assert!(catch_unwind(AssertUnwindSafe(|| sender.push(1))).is_err());
742        assert!(catch_unwind(AssertUnwindSafe(|| stream.try_next())).is_err());
743        assert!(catch_unwind(AssertUnwindSafe(|| pollster::block_on(stream.next()))).is_err());
744    }
745
746    #[derive(Default)]
747    struct CountingWake(AtomicUsize);
748
749    impl Wake for CountingWake {
750        fn wake(self: Arc<Self>) {
751            self.0.fetch_add(1, Ordering::SeqCst);
752        }
753    }
754
755    fn counting_waker() -> (Arc<CountingWake>, Waker) {
756        let probe = Arc::new(CountingWake::default());
757        let waker = Waker::from(Arc::clone(&probe));
758        (probe, waker)
759    }
760
761    fn poll_once<F: Future + Unpin>(future: &mut F, waker: &Waker) -> Poll<F::Output> {
762        Pin::new(future).poll(&mut Context::from_waker(waker))
763    }
764
765    #[test]
766    fn every_waiting_consumer_is_woken() {
767        let (stream, sender) = BoundedAsyncStream::new(4);
768        let (first_probe, first_waker) = counting_waker();
769        let (second_probe, second_waker) = counting_waker();
770        let mut first = stream.next();
771        let mut second = stream.next();
772
773        assert_eq!(poll_once(&mut first, &first_waker), Poll::Pending);
774        assert_eq!(poll_once(&mut second, &second_waker), Poll::Pending);
775
776        sender.push(1);
777        assert_eq!(first_probe.0.load(Ordering::SeqCst), 1);
778        assert_eq!(second_probe.0.load(Ordering::SeqCst), 1);
779
780        assert_eq!(poll_once(&mut second, &second_waker), Poll::Ready(Some(1)));
781        assert_eq!(poll_once(&mut first, &first_waker), Poll::Pending);
782
783        drop(sender);
784        assert_eq!(first_probe.0.load(Ordering::SeqCst), 2);
785        assert_eq!(second_probe.0.load(Ordering::SeqCst), 1);
786        assert_eq!(poll_once(&mut first, &first_waker), Poll::Ready(None));
787    }
788
789    #[test]
790    fn dropped_next_future_releases_its_waker() {
791        let (stream, sender) = BoundedAsyncStream::<u32>::new(1);
792        let (first_probe, first_waker) = counting_waker();
793        let (second_probe, second_waker) = counting_waker();
794        let mut first = stream.next();
795        let mut second = stream.next();
796
797        for _ in 0..3 {
798            assert_eq!(poll_once(&mut first, &first_waker), Poll::Pending);
799        }
800        assert_eq!(poll_once(&mut second, &second_waker), Poll::Pending);
801        assert_eq!(stream.shared.lock_state().waiters.len(), 2);
802
803        drop(first);
804        assert_eq!(stream.shared.lock_state().waiters.len(), 1);
805        drop(first_waker);
806        assert_eq!(Arc::strong_count(&first_probe), 1);
807
808        sender.push(5);
809        assert_eq!(first_probe.0.load(Ordering::SeqCst), 0);
810        assert_eq!(second_probe.0.load(Ordering::SeqCst), 1);
811        assert_eq!(poll_once(&mut second, &second_waker), Poll::Ready(Some(5)));
812        drop(second);
813        assert!(stream.shared.lock_state().waiters.is_empty());
814    }
815
816    #[test]
817    fn concurrent_consumers_drain_every_item() {
818        const CONSUMERS: usize = 4;
819        const ITEMS: usize = 2_000;
820
821        let (stream, sender) = BoundedAsyncStream::<usize>::new(4);
822        let stream = Arc::new(stream);
823        let (done_tx, done_rx) = mpsc::channel();
824        let consumers = (0..CONSUMERS)
825            .map(|_| {
826                let stream = Arc::clone(&stream);
827                let done_tx = done_tx.clone();
828                thread::spawn(move || {
829                    let mut received = 0;
830                    while pollster::block_on(stream.next()).is_some() {
831                        received += 1;
832                    }
833                    done_tx.send(received).unwrap();
834                })
835            })
836            .collect::<Vec<_>>();
837
838        for item in 0..ITEMS {
839            assert_eq!(sender.push_or_block(item), Ok(()));
840        }
841        drop(sender);
842
843        let received: usize = (0..CONSUMERS)
844            .map(|_| done_rx.recv_timeout(TEST_TIMEOUT).expect("a consumer hung"))
845            .sum();
846        assert_eq!(received, ITEMS);
847        for consumer in consumers {
848            consumer.join().unwrap();
849        }
850    }
851
852    #[cfg(feature = "futures-stream")]
853    #[test]
854    fn stream_poll_keeps_a_single_registration() {
855        let (mut stream, sender) = BoundedAsyncStream::<u32>::new(1);
856        let (first_probe, first_waker) = counting_waker();
857        let (second_probe, second_waker) = counting_waker();
858
859        let mut first_cx = Context::from_waker(&first_waker);
860        assert_eq!(
861            Pin::new(&mut stream).poll_next(&mut first_cx),
862            Poll::Pending
863        );
864        let mut second_cx = Context::from_waker(&second_waker);
865        assert_eq!(
866            Pin::new(&mut stream).poll_next(&mut second_cx),
867            Poll::Pending
868        );
869        assert_eq!(stream.shared.lock_state().waiters.len(), 1);
870
871        sender.push(3);
872        assert_eq!(first_probe.0.load(Ordering::SeqCst), 0);
873        assert_eq!(second_probe.0.load(Ordering::SeqCst), 1);
874        assert_eq!(
875            Pin::new(&mut stream).poll_next(&mut second_cx),
876            Poll::Ready(Some(3))
877        );
878    }
879}