diff --git a/tokio-stream/src/wrappers/mpsc_bounded.rs b/tokio-stream/src/wrappers/mpsc_bounded.rs index 34b2e020d..d2c495110 100644 --- a/tokio-stream/src/wrappers/mpsc_bounded.rs +++ b/tokio-stream/src/wrappers/mpsc_bounded.rs @@ -67,6 +67,25 @@ impl Stream for ReceiverStream { fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { self.inner.poll_recv(cx) } + + /// Returns the bounds of the stream based on the underlying receiver. + /// + /// For open channels, it returns `(receiver.len(), None)`. + /// + /// For closed channels, it returns `(receiver.len(), Some(used_capacity))` + /// where `used_capacity` is calculated as `receiver.max_capacity() - + /// receiver.capacity()`. This accounts for any [`Permit`] that is still + /// able to send a message. + /// + /// [`Permit`]: struct@tokio::sync::mpsc::Permit + fn size_hint(&self) -> (usize, Option) { + if self.inner.is_closed() { + let used_capacity = self.inner.max_capacity() - self.inner.capacity(); + (self.inner.len(), Some(used_capacity)) + } else { + (self.inner.len(), None) + } + } } impl AsRef> for ReceiverStream { diff --git a/tokio-stream/src/wrappers/mpsc_unbounded.rs b/tokio-stream/src/wrappers/mpsc_unbounded.rs index 904c98bef..6221203b5 100644 --- a/tokio-stream/src/wrappers/mpsc_unbounded.rs +++ b/tokio-stream/src/wrappers/mpsc_unbounded.rs @@ -61,6 +61,20 @@ impl Stream for UnboundedReceiverStream { fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll> { self.inner.poll_recv(cx) } + + /// Returns the bounds of the stream based on the underlying receiver. + /// + /// For open channels, it returns `(receiver.len(), None)`. + /// + /// For closed channels, it returns `(receiver.len(), receiver.len())`. + fn size_hint(&self) -> (usize, Option) { + if self.inner.is_closed() { + let len = self.inner.len(); + (len, Some(len)) + } else { + (self.inner.len(), None) + } + } } impl AsRef> for UnboundedReceiverStream { diff --git a/tokio-stream/tests/mpsc_bounded_stream.rs b/tokio-stream/tests/mpsc_bounded_stream.rs new file mode 100644 index 000000000..160c5a18e --- /dev/null +++ b/tokio-stream/tests/mpsc_bounded_stream.rs @@ -0,0 +1,109 @@ +use futures::{Stream, StreamExt}; +use tokio::sync::mpsc; +use tokio_stream::wrappers::ReceiverStream; + +#[tokio::test] +async fn size_hint_stream_open() { + let (tx, rx) = mpsc::channel(4); + + tx.send(1).await.unwrap(); + tx.send(2).await.unwrap(); + + let mut stream = ReceiverStream::new(rx); + + assert_eq!(stream.size_hint(), (2, None)); + stream.next().await; + assert_eq!(stream.size_hint(), (1, None)); + stream.next().await; + assert_eq!(stream.size_hint(), (0, None)); +} + +#[tokio::test] +async fn size_hint_stream_closed() { + let (tx, rx) = mpsc::channel(4); + + tx.send(1).await.unwrap(); + tx.send(2).await.unwrap(); + + let mut stream = ReceiverStream::new(rx); + stream.close(); + + assert_eq!(stream.size_hint(), (2, Some(2))); + stream.next().await; + assert_eq!(stream.size_hint(), (1, Some(1))); + stream.next().await; + assert_eq!(stream.size_hint(), (0, Some(0))); +} + +#[tokio::test] +async fn size_hint_sender_dropped() { + let (tx, rx) = mpsc::channel(4); + + tx.send(1).await.unwrap(); + tx.send(2).await.unwrap(); + + let mut stream = ReceiverStream::new(rx); + drop(tx); + + assert_eq!(stream.size_hint(), (2, Some(2))); + stream.next().await; + assert_eq!(stream.size_hint(), (1, Some(1))); + stream.next().await; + assert_eq!(stream.size_hint(), (0, Some(0))); +} + +#[test] +fn size_hint_stream_instantly_closed() { + let (_tx, rx) = mpsc::channel::(4); + + let mut stream = ReceiverStream::new(rx); + stream.close(); + + assert_eq!(stream.size_hint(), (0, Some(0))); +} + +#[tokio::test] +async fn size_hint_stream_closed_permits_send() { + let (tx, rx) = mpsc::channel(4); + + tx.send(1).await.unwrap(); + let permit1 = tx.reserve().await.unwrap(); + let permit2 = tx.reserve().await.unwrap(); + + let mut stream = ReceiverStream::new(rx); + stream.close(); + + assert_eq!(stream.size_hint(), (1, Some(3))); + permit1.send(2); + assert_eq!(stream.size_hint(), (2, Some(3))); + stream.next().await; + assert_eq!(stream.size_hint(), (1, Some(2))); + stream.next().await; + assert_eq!(stream.size_hint(), (0, Some(1))); + permit2.send(3); + assert_eq!(stream.size_hint(), (1, Some(1))); + stream.next().await; + assert_eq!(stream.size_hint(), (0, Some(0))); + assert_eq!(stream.next().await, None); +} + +#[tokio::test] +async fn size_hint_stream_closed_permits_drop() { + let (tx, rx) = mpsc::channel(4); + + tx.send(1).await.unwrap(); + let permit1 = tx.reserve().await.unwrap(); + let permit2 = tx.reserve().await.unwrap(); + + let mut stream = ReceiverStream::new(rx); + stream.close(); + + assert_eq!(stream.size_hint(), (1, Some(3))); + drop(permit1); + assert_eq!(stream.size_hint(), (1, Some(2))); + stream.next().await; + assert_eq!(stream.size_hint(), (0, Some(1))); + drop(permit2); + assert_eq!(stream.size_hint(), (0, Some(0))); + assert_eq!(stream.next().await, None); +} diff --git a/tokio-stream/tests/mpsc_unbounded_stream.rs b/tokio-stream/tests/mpsc_unbounded_stream.rs new file mode 100644 index 000000000..c3a02c836 --- /dev/null +++ b/tokio-stream/tests/mpsc_unbounded_stream.rs @@ -0,0 +1,63 @@ +use futures::{Stream, StreamExt}; +use tokio::sync::mpsc; +use tokio_stream::wrappers::UnboundedReceiverStream; + +#[tokio::test] +async fn size_hint_stream_open() { + let (tx, rx) = mpsc::unbounded_channel(); + + tx.send(1).unwrap(); + tx.send(2).unwrap(); + + let mut stream = UnboundedReceiverStream::new(rx); + + assert_eq!(stream.size_hint(), (2, None)); + stream.next().await; + assert_eq!(stream.size_hint(), (1, None)); + stream.next().await; + assert_eq!(stream.size_hint(), (0, None)); +} + +#[tokio::test] +async fn size_hint_stream_closed() { + let (tx, rx) = mpsc::unbounded_channel(); + + tx.send(1).unwrap(); + tx.send(2).unwrap(); + + let mut stream = UnboundedReceiverStream::new(rx); + stream.close(); + + assert_eq!(stream.size_hint(), (2, Some(2))); + stream.next().await; + assert_eq!(stream.size_hint(), (1, Some(1))); + stream.next().await; + assert_eq!(stream.size_hint(), (0, Some(0))); +} + +#[tokio::test] +async fn size_hint_sender_dropped() { + let (tx, rx) = mpsc::unbounded_channel(); + + tx.send(1).unwrap(); + tx.send(2).unwrap(); + + let mut stream = UnboundedReceiverStream::new(rx); + drop(tx); + + assert_eq!(stream.size_hint(), (2, Some(2))); + stream.next().await; + assert_eq!(stream.size_hint(), (1, Some(1))); + stream.next().await; + assert_eq!(stream.size_hint(), (0, Some(0))); +} + +#[test] +fn size_hint_stream_instantly_closed() { + let (_tx, rx) = mpsc::unbounded_channel::(); + + let mut stream = UnboundedReceiverStream::new(rx); + stream.close(); + + assert_eq!(stream.size_hint(), (0, Some(0))); +}