diff --git a/tokio-stream/src/stream_ext/chain.rs b/tokio-stream/src/stream_ext/chain.rs index f3d360d53..2ddf21ec9 100644 --- a/tokio-stream/src/stream_ext/chain.rs +++ b/tokio-stream/src/stream_ext/chain.rs @@ -3,6 +3,7 @@ use crate::Stream; use core::pin::Pin; use core::task::{ready, Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -48,3 +49,13 @@ where super::merge_size_hints(self.a.size_hint(), self.b.size_hint()) } } + +impl FusedStream for Chain +where + T: Stream, + U: FusedStream, +{ + fn is_terminated(&self) -> bool { + self.a.is_terminated() && self.b.is_terminated() + } +} diff --git a/tokio-stream/src/stream_ext/filter.rs b/tokio-stream/src/stream_ext/filter.rs index 1d5defb19..f4a5eeef4 100644 --- a/tokio-stream/src/stream_ext/filter.rs +++ b/tokio-stream/src/stream_ext/filter.rs @@ -3,6 +3,7 @@ use crate::Stream; use core::fmt; use core::pin::Pin; use core::task::{ready, Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -56,3 +57,13 @@ where (0, self.stream.size_hint().1) // can't know a lower bound, due to the predicate } } + +impl FusedStream for Filter +where + St: FusedStream, + F: FnMut(&St::Item) -> bool, +{ + fn is_terminated(&self) -> bool { + self.stream.is_terminated() + } +} diff --git a/tokio-stream/src/stream_ext/filter_map.rs b/tokio-stream/src/stream_ext/filter_map.rs index 6658d71f0..01244860c 100644 --- a/tokio-stream/src/stream_ext/filter_map.rs +++ b/tokio-stream/src/stream_ext/filter_map.rs @@ -3,6 +3,7 @@ use crate::Stream; use core::fmt; use core::pin::Pin; use core::task::{ready, Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -56,3 +57,13 @@ where (0, self.stream.size_hint().1) // can't know a lower bound, due to the predicate } } + +impl FusedStream for FilterMap +where + St: FusedStream, + F: FnMut(St::Item) -> Option, +{ + fn is_terminated(&self) -> bool { + self.stream.is_terminated() + } +} diff --git a/tokio-stream/src/stream_ext/map.rs b/tokio-stream/src/stream_ext/map.rs index e6b47cd25..c95603af4 100644 --- a/tokio-stream/src/stream_ext/map.rs +++ b/tokio-stream/src/stream_ext/map.rs @@ -3,6 +3,7 @@ use crate::Stream; use core::fmt; use core::pin::Pin; use core::task::{Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -49,3 +50,13 @@ where self.stream.size_hint() } } + +impl FusedStream for Map +where + St: FusedStream, + F: FnMut(St::Item) -> T, +{ + fn is_terminated(&self) -> bool { + self.stream.is_terminated() + } +} diff --git a/tokio-stream/src/stream_ext/merge.rs b/tokio-stream/src/stream_ext/merge.rs index 4f7022b1f..2f4c3e650 100644 --- a/tokio-stream/src/stream_ext/merge.rs +++ b/tokio-stream/src/stream_ext/merge.rs @@ -3,6 +3,7 @@ use crate::Stream; use core::pin::Pin; use core::task::{Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -57,6 +58,16 @@ where } } +impl FusedStream for Merge +where + T: Stream, + U: Stream, +{ + fn is_terminated(&self) -> bool { + self.a.is_terminated() && self.b.is_terminated() + } +} + fn poll_next( first: Pin<&mut T>, second: Pin<&mut U>, diff --git a/tokio-stream/src/stream_ext/skip.rs b/tokio-stream/src/stream_ext/skip.rs index dd310b856..0e052f14c 100644 --- a/tokio-stream/src/stream_ext/skip.rs +++ b/tokio-stream/src/stream_ext/skip.rs @@ -3,6 +3,7 @@ use crate::Stream; use core::fmt; use core::pin::Pin; use core::task::{ready, Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -61,3 +62,12 @@ where (lower, upper) } } + +impl FusedStream for Skip +where + St: FusedStream, +{ + fn is_terminated(&self) -> bool { + self.stream.is_terminated() + } +} diff --git a/tokio-stream/src/stream_ext/skip_while.rs b/tokio-stream/src/stream_ext/skip_while.rs index d1accd529..03d2a0d2f 100644 --- a/tokio-stream/src/stream_ext/skip_while.rs +++ b/tokio-stream/src/stream_ext/skip_while.rs @@ -3,6 +3,7 @@ use crate::Stream; use core::fmt; use core::pin::Pin; use core::task::{ready, Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -71,3 +72,13 @@ where (lower, upper) } } + +impl FusedStream for SkipWhile +where + St: FusedStream, + F: FnMut(&St::Item) -> bool, +{ + fn is_terminated(&self) -> bool { + self.stream.is_terminated() + } +} diff --git a/tokio-stream/src/stream_ext/take.rs b/tokio-stream/src/stream_ext/take.rs index 07b5c3d57..3bd08ca0e 100644 --- a/tokio-stream/src/stream_ext/take.rs +++ b/tokio-stream/src/stream_ext/take.rs @@ -4,6 +4,7 @@ use core::cmp; use core::fmt; use core::pin::Pin; use core::task::{Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -74,3 +75,12 @@ where (lower, upper) } } + +impl FusedStream for Take +where + St: Stream, +{ + fn is_terminated(&self) -> bool { + self.remaining == 0 + } +} diff --git a/tokio-stream/src/stream_ext/take_while.rs b/tokio-stream/src/stream_ext/take_while.rs index 5ce4dd98a..eee104059 100644 --- a/tokio-stream/src/stream_ext/take_while.rs +++ b/tokio-stream/src/stream_ext/take_while.rs @@ -3,6 +3,7 @@ use crate::Stream; use core::fmt; use core::pin::Pin; use core::task::{Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -77,3 +78,13 @@ where (0, upper) } } + +impl FusedStream for TakeWhile +where + St: Stream, + F: FnMut(&St::Item) -> bool, +{ + fn is_terminated(&self) -> bool { + self.done + } +} diff --git a/tokio-stream/src/stream_ext/then.rs b/tokio-stream/src/stream_ext/then.rs index cc7caa721..c4314e6d0 100644 --- a/tokio-stream/src/stream_ext/then.rs +++ b/tokio-stream/src/stream_ext/then.rs @@ -4,6 +4,7 @@ use core::fmt; use core::future::Future; use core::pin::Pin; use core::task::{Context, Poll}; +use futures_core::FusedStream; use pin_project_lite::pin_project; pin_project! { @@ -81,3 +82,14 @@ where (lower, upper) } } + +impl FusedStream for Then +where + St: FusedStream, + Fut: Future, + F: FnMut(St::Item) -> Fut, +{ + fn is_terminated(&self) -> bool { + self.future.is_none() && self.stream.is_terminated() + } +} diff --git a/tokio-stream/tests/stream_fused.rs b/tokio-stream/tests/stream_fused.rs new file mode 100644 index 000000000..5eb7166ef --- /dev/null +++ b/tokio-stream/tests/stream_fused.rs @@ -0,0 +1,188 @@ +use futures_core::FusedStream; +use tokio_stream::StreamExt; + +// Helper: a fused base stream built from a vec +fn fused_iter(items: Vec) -> impl FusedStream { + tokio_stream::iter(items).fuse() +} + +// ── map ────────────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn map_not_terminated_before_done() { + let stream = fused_iter(vec![1, 2]).map(|x| x * 2); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn map_terminated_after_inner_done() { + let mut stream = fused_iter(vec![1]).map(|x| x * 2); + assert_eq!(stream.next().await, Some(2)); + assert_eq!(stream.next().await, None); + assert!(stream.is_terminated()); +} + +// ── filter ─────────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn filter_not_terminated_before_done() { + let stream = fused_iter(vec![1, 2]).filter(|x| *x > 0); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn filter_terminated_after_inner_done() { + let mut stream = fused_iter(vec![1]).filter(|x| *x > 0); + assert_eq!(stream.next().await, Some(1)); + assert_eq!(stream.next().await, None); + assert!(stream.is_terminated()); +} + +// ── filter_map ─────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn filter_map_not_terminated_before_done() { + let stream = fused_iter(vec![1, 2]).filter_map(Some); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn filter_map_terminated_after_inner_done() { + let mut stream = fused_iter(vec![1]).filter_map(|x| Some(x * 10)); + assert_eq!(stream.next().await, Some(10)); + assert_eq!(stream.next().await, None); + assert!(stream.is_terminated()); +} + +// ── skip ───────────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn skip_not_terminated_before_done() { + let stream = fused_iter(vec![1, 2, 3]).skip(1); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn skip_terminated_after_inner_done() { + let mut stream = fused_iter(vec![1, 2]).skip(1); + assert_eq!(stream.next().await, Some(2)); + assert_eq!(stream.next().await, None); + assert!(stream.is_terminated()); +} + +// ── skip_while ─────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn skip_while_not_terminated_before_done() { + let stream = fused_iter(vec![1, 2, 3]).skip_while(|x| *x < 2); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn skip_while_terminated_after_inner_done() { + let mut stream = fused_iter(vec![1, 2]).skip_while(|x| *x < 2); + assert_eq!(stream.next().await, Some(2)); + assert_eq!(stream.next().await, None); + assert!(stream.is_terminated()); +} + +// ── take ───────────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn take_not_terminated_before_limit() { + let stream = tokio_stream::iter(vec![1, 2, 3]).take(2); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn take_terminated_when_remaining_zero() { + let mut stream = tokio_stream::iter(vec![1, 2]).take(2); + assert_eq!(stream.next().await, Some(1)); + assert!(!stream.is_terminated()); + assert_eq!(stream.next().await, Some(2)); + // remaining hits 0 after getting the second item + assert!(stream.is_terminated()); + assert_eq!(stream.next().await, None); +} + +#[tokio::test] +async fn take_zero_is_immediately_terminated() { + let stream = tokio_stream::iter(vec![1, 2]).take(0); + assert!(stream.is_terminated()); +} + +// ── take_while ─────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn take_while_not_terminated_before_predicate_fails() { + let stream = tokio_stream::iter(vec![1, 2, 3]).take_while(|x| *x < 10); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn take_while_terminated_after_predicate_fails() { + let mut stream = tokio_stream::iter(vec![1, 5, 2]).take_while(|x| *x < 3); + assert_eq!(stream.next().await, Some(1)); + assert!(!stream.is_terminated()); + // predicate fails on 5 → done flag set + assert_eq!(stream.next().await, None); + assert!(stream.is_terminated()); +} + +// ── then ───────────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn then_not_terminated_before_done() { + let stream = fused_iter(vec![1, 2]).then(|x| async move { x * 2 }); + tokio::pin!(stream); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn then_terminated_after_inner_done_and_no_pending_future() { + let stream = fused_iter(vec![1]).then(|x| async move { x * 2 }); + tokio::pin!(stream); + assert_eq!(stream.next().await, Some(2)); + assert_eq!(stream.next().await, None); + // inner stream done AND no in-flight future + assert!(stream.is_terminated()); +} + +// ── chain ──────────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn chain_not_terminated_while_either_has_items() { + let stream = fused_iter(vec![1]).chain(fused_iter(vec![2])); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn chain_terminated_only_after_both_done() { + let mut stream = fused_iter(vec![1]).chain(fused_iter(vec![2])); + assert_eq!(stream.next().await, Some(1)); + assert!(!stream.is_terminated()); // b still has items + assert_eq!(stream.next().await, Some(2)); + assert_eq!(stream.next().await, None); + assert!(stream.is_terminated()); // both done now +} + +// ── merge ──────────────────────────────────────────────────────────────────── + +#[tokio::test] +async fn merge_not_terminated_while_either_has_items() { + let stream = fused_iter(vec![1]).merge(fused_iter(vec![2])); + assert!(!stream.is_terminated()); +} + +#[tokio::test] +async fn merge_terminated_only_after_both_done() { + let mut stream = fused_iter(vec![1]).merge(fused_iter(vec![2])); + // drain both + let mut collected = vec![]; + while let Some(x) = stream.next().await { + collected.push(x); + } + assert_eq!(stream.next().await, None); + assert!(stream.is_terminated()); + assert_eq!(collected.len(), 2); +}