diff --git a/tokio-stream/src/lib.rs b/tokio-stream/src/lib.rs index b6e651c7b..6ff1085a5 100644 --- a/tokio-stream/src/lib.rs +++ b/tokio-stream/src/lib.rs @@ -73,6 +73,9 @@ #[macro_use] mod macros; +mod poll_fn; +pub(crate) use poll_fn::poll_fn; + pub mod wrappers; mod stream_ext; diff --git a/tokio-stream/src/poll_fn.rs b/tokio-stream/src/poll_fn.rs new file mode 100644 index 000000000..744f22f02 --- /dev/null +++ b/tokio-stream/src/poll_fn.rs @@ -0,0 +1,35 @@ +use std::future::Future; +use std::pin::Pin; +use std::task::{Context, Poll}; + +pub(crate) struct PollFn { + f: F, +} + +pub(crate) fn poll_fn(f: F) -> PollFn +where + F: FnMut(&mut Context<'_>) -> Poll, +{ + PollFn { f } +} + +impl Future for PollFn +where + F: FnMut(&mut Context<'_>) -> Poll, +{ + type Output = T; + + fn poll(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll { + // Safety: We never construct a `Pin<&mut F>` anywhere, so accessing `f` + // mutably in an unpinned way is sound. + // + // This use of unsafe cannot be replaced with the pin-project macro + // because: + // * If we put `#[pin]` on the field, then it gives us a `Pin<&mut F>`, + // which we can't use to call the closure. + // * If we don't put `#[pin]` on the field, then it makes `PollFn` be + // unconditionally `Unpin`, which we also don't want. + let me = unsafe { Pin::into_inner_unchecked(self) }; + (me.f)(cx) + } +} diff --git a/tokio-stream/src/stream_map.rs b/tokio-stream/src/stream_map.rs index 3f424eca2..41ab9648c 100644 --- a/tokio-stream/src/stream_map.rs +++ b/tokio-stream/src/stream_map.rs @@ -1,4 +1,4 @@ -use crate::Stream; +use crate::{poll_fn, Stream}; use std::borrow::Borrow; use std::hash::Hash; @@ -561,6 +561,110 @@ impl Default for StreamMap { } } +impl StreamMap +where + K: Clone + Unpin, + V: Stream + Unpin, +{ + /// Receives multiple items on this [`StreamMap`], extending the provided `buffer`. + /// + /// This method returns the number of items that is appended to the `buffer`. + /// + /// Note that this method does not guarantee that exactly `limit` items + /// are received. Rather, if at least one item is available, it returns + /// as many items as it can up to the given limit. This method returns + /// zero only if the `StreamMap` is empty (or if `limit` is zero). + /// + /// # Cancel safety + /// + /// This method is cancel safe. If `next_many` is used as the event in a + /// [`tokio::select!`](tokio::select) statement and some other branch + /// completes first, it is guaranteed that no items were received on any of + /// the underlying streams. + pub async fn next_many(&mut self, buffer: &mut Vec<(K, V::Item)>, limit: usize) -> usize { + poll_fn(|cx| self.poll_next_many(cx, buffer, limit)).await + } + + /// Polls to receive multiple items on this `StreamMap`, extending the provided `buffer`. + /// + /// This method returns: + /// * `Poll::Pending` if no items are available but the `StreamMap` is not empty. + /// * `Poll::Ready(count)` where `count` is the number of items successfully received and + /// stored in `buffer`. This can be less than, or equal to, `limit`. + /// * `Poll::Ready(0)` if `limit` is set to zero or when the `StreamMap` is empty. + /// + /// Note that this method does not guarantee that exactly `limit` items + /// are received. Rather, if at least one item is available, it returns + /// as many items as it can up to the given limit. This method returns + /// zero only if the `StreamMap` is empty (or if `limit` is zero). + pub fn poll_next_many( + &mut self, + cx: &mut Context<'_>, + buffer: &mut Vec<(K, V::Item)>, + limit: usize, + ) -> Poll { + if limit == 0 || self.entries.is_empty() { + return Poll::Ready(0); + } + + let mut added = 0; + + let start = self::rand::thread_rng_n(self.entries.len() as u32) as usize; + let mut idx = start; + + while added < limit { + // Indicates whether at least one stream returned a value when polled or not + let mut should_loop = false; + + for _ in 0..self.entries.len() { + let (_, stream) = &mut self.entries[idx]; + + match Pin::new(stream).poll_next(cx) { + Poll::Ready(Some(val)) => { + added += 1; + + let key = self.entries[idx].0.clone(); + buffer.push((key, val)); + + should_loop = true; + + idx = idx.wrapping_add(1) % self.entries.len(); + } + Poll::Ready(None) => { + // Remove the entry + self.entries.swap_remove(idx); + + // Check if this was the last entry, if so the cursor needs + // to wrap + if idx == self.entries.len() { + idx = 0; + } else if idx < start && start <= self.entries.len() { + // The stream being swapped into the current index has + // already been polled, so skip it. + idx = idx.wrapping_add(1) % self.entries.len(); + } + } + Poll::Pending => { + idx = idx.wrapping_add(1) % self.entries.len(); + } + } + } + + if !should_loop { + break; + } + } + + if added > 0 { + Poll::Ready(added) + } else if self.entries.is_empty() { + Poll::Ready(0) + } else { + Poll::Pending + } + } +} + impl Stream for StreamMap where K: Clone + Unpin, diff --git a/tokio-stream/tests/stream_stream_map.rs b/tokio-stream/tests/stream_stream_map.rs index b6b87e9d0..5acceb5c9 100644 --- a/tokio-stream/tests/stream_stream_map.rs +++ b/tokio-stream/tests/stream_stream_map.rs @@ -1,14 +1,17 @@ +use futures::stream::iter; use tokio_stream::{self as stream, pending, Stream, StreamExt, StreamMap}; use tokio_test::{assert_ok, assert_pending, assert_ready, task}; +use std::future::{poll_fn, Future}; +use std::pin::{pin, Pin}; +use std::task::Poll; + mod support { pub(crate) mod mpsc; } use support::mpsc; -use std::pin::Pin; - macro_rules! assert_ready_some { ($($t:tt)*) => { match assert_ready!($($t)*) { @@ -328,3 +331,233 @@ fn one_ready_many_none() { fn pin_box + 'static, U>(s: T) -> Pin>> { Box::pin(s) } + +type UsizeStream = Pin + Send>>; + +#[tokio::test] +async fn poll_next_many_zero() { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(pending()) as UsizeStream); + + let n = poll_fn(|cx| stream_map.poll_next_many(cx, &mut vec![], 0)).await; + + assert_eq!(n, 0); +} + +#[tokio::test] +async fn poll_next_many_empty() { + let mut stream_map: StreamMap = StreamMap::new(); + + let n = poll_fn(|cx| stream_map.poll_next_many(cx, &mut vec![], 1)).await; + + assert_eq!(n, 0); +} + +#[tokio::test] +async fn poll_next_many_pending() { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(pending()) as UsizeStream); + + let mut is_pending = false; + poll_fn(|cx| { + let poll = stream_map.poll_next_many(cx, &mut vec![], 1); + + is_pending = poll.is_pending(); + + Poll::Ready(()) + }) + .await; + + assert!(is_pending); +} + +#[tokio::test] +async fn poll_next_many_not_enough() { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(iter([0usize].into_iter())) as UsizeStream); + stream_map.insert(1, Box::pin(iter([1usize].into_iter())) as UsizeStream); + + let mut buffer = vec![]; + let n = poll_fn(|cx| stream_map.poll_next_many(cx, &mut buffer, 3)).await; + + assert_eq!(n, 2); + assert_eq!(buffer.len(), 2); + assert!(buffer.contains(&(0, 0))); + assert!(buffer.contains(&(1, 1))); +} + +#[tokio::test] +async fn poll_next_many_enough() { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(iter([0usize].into_iter())) as UsizeStream); + stream_map.insert(1, Box::pin(iter([1usize].into_iter())) as UsizeStream); + + let mut buffer = vec![]; + let n = poll_fn(|cx| stream_map.poll_next_many(cx, &mut buffer, 2)).await; + + assert_eq!(n, 2); + assert_eq!(buffer.len(), 2); + assert!(buffer.contains(&(0, 0))); + assert!(buffer.contains(&(1, 1))); +} + +#[tokio::test] +async fn poll_next_many_correctly_loops_around() { + for _ in 0..10 { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(iter([0usize].into_iter())) as UsizeStream); + stream_map.insert(1, Box::pin(iter([0usize, 1].into_iter())) as UsizeStream); + stream_map.insert(2, Box::pin(iter([0usize, 1, 2].into_iter())) as UsizeStream); + + let mut buffer = vec![]; + + let n = poll_fn(|cx| stream_map.poll_next_many(cx, &mut buffer, 3)).await; + assert_eq!(n, 3); + assert_eq!( + std::mem::take(&mut buffer) + .into_iter() + .map(|(_, v)| v) + .collect::>(), + vec![0, 0, 0] + ); + + let n = poll_fn(|cx| stream_map.poll_next_many(cx, &mut buffer, 2)).await; + assert_eq!(n, 2); + assert_eq!( + std::mem::take(&mut buffer) + .into_iter() + .map(|(_, v)| v) + .collect::>(), + vec![1, 1] + ); + + let n = poll_fn(|cx| stream_map.poll_next_many(cx, &mut buffer, 1)).await; + assert_eq!(n, 1); + assert_eq!( + std::mem::take(&mut buffer) + .into_iter() + .map(|(_, v)| v) + .collect::>(), + vec![2] + ); + } +} + +#[tokio::test] +async fn next_many_zero() { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(pending()) as UsizeStream); + + let n = poll_fn(|cx| pin!(stream_map.next_many(&mut vec![], 0)).poll(cx)).await; + + assert_eq!(n, 0); +} + +#[tokio::test] +async fn next_many_empty() { + let mut stream_map: StreamMap = StreamMap::new(); + + let n = stream_map.next_many(&mut vec![], 1).await; + + assert_eq!(n, 0); +} + +#[tokio::test] +async fn next_many_pending() { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(pending()) as UsizeStream); + + let mut is_pending = false; + poll_fn(|cx| { + let poll = pin!(stream_map.next_many(&mut vec![], 1)).poll(cx); + + is_pending = poll.is_pending(); + + Poll::Ready(()) + }) + .await; + + assert!(is_pending); +} + +#[tokio::test] +async fn next_many_not_enough() { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(iter([0usize].into_iter())) as UsizeStream); + stream_map.insert(1, Box::pin(iter([1usize].into_iter())) as UsizeStream); + + let mut buffer = vec![]; + let n = poll_fn(|cx| pin!(stream_map.next_many(&mut buffer, 3)).poll(cx)).await; + + assert_eq!(n, 2); + assert_eq!(buffer.len(), 2); + assert!(buffer.contains(&(0, 0))); + assert!(buffer.contains(&(1, 1))); +} + +#[tokio::test] +async fn next_many_enough() { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(iter([0usize].into_iter())) as UsizeStream); + stream_map.insert(1, Box::pin(iter([1usize].into_iter())) as UsizeStream); + + let mut buffer = vec![]; + let n = poll_fn(|cx| pin!(stream_map.next_many(&mut buffer, 2)).poll(cx)).await; + + assert_eq!(n, 2); + assert_eq!(buffer.len(), 2); + assert!(buffer.contains(&(0, 0))); + assert!(buffer.contains(&(1, 1))); +} + +#[tokio::test] +async fn next_many_correctly_loops_around() { + for _ in 0..10 { + let mut stream_map: StreamMap = StreamMap::new(); + + stream_map.insert(0, Box::pin(iter([0usize].into_iter())) as UsizeStream); + stream_map.insert(1, Box::pin(iter([0usize, 1].into_iter())) as UsizeStream); + stream_map.insert(2, Box::pin(iter([0usize, 1, 2].into_iter())) as UsizeStream); + + let mut buffer = vec![]; + + let n = poll_fn(|cx| pin!(stream_map.next_many(&mut buffer, 3)).poll(cx)).await; + assert_eq!(n, 3); + assert_eq!( + std::mem::take(&mut buffer) + .into_iter() + .map(|(_, v)| v) + .collect::>(), + vec![0, 0, 0] + ); + + let n = poll_fn(|cx| pin!(stream_map.next_many(&mut buffer, 2)).poll(cx)).await; + assert_eq!(n, 2); + assert_eq!( + std::mem::take(&mut buffer) + .into_iter() + .map(|(_, v)| v) + .collect::>(), + vec![1, 1] + ); + + let n = poll_fn(|cx| pin!(stream_map.next_many(&mut buffer, 1)).poll(cx)).await; + assert_eq!(n, 1); + assert_eq!( + std::mem::take(&mut buffer) + .into_iter() + .map(|(_, v)| v) + .collect::>(), + vec![2] + ); + } +}