use alloc::collections::BTreeSet;
use core::fmt::{self, Debug};
use core::ops::{Deref, DerefMut};
use core::pin::Pin;
use core::task::{Context, Poll};
use futures_core::Stream;
use smallvec::{smallvec, SmallVec};
use crate::utils::{ChunkedVec, PollState, PollVec, WakerVec};
#[must_use = "`StreamGroup` does nothing if not iterated over"]
#[derive(Default)]
#[pin_project::pin_project]
pub struct StreamGroup<S> {
#[pin]
streams: ChunkedVec<S>,
wakers: WakerVec,
states: PollVec,
keys: BTreeSet<usize>,
key_removal_queue: SmallVec<[usize; 10]>,
}
impl<T: Debug> Debug for StreamGroup<T> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("StreamGroup")
.field("streams", &"[..]")
.field("len", &self.len())
.field("capacity", &self.capacity())
.finish()
}
}
impl<S> StreamGroup<S> {
pub fn new() -> Self {
Self::with_capacity(0)
}
pub fn with_capacity(capacity: usize) -> Self {
Self {
streams: ChunkedVec::with_capacity(capacity),
wakers: WakerVec::new(capacity),
states: PollVec::new(capacity),
keys: BTreeSet::new(),
key_removal_queue: smallvec![],
}
}
#[inline(always)]
pub fn len(&self) -> usize {
self.streams.len()
}
pub fn capacity(&self) -> usize {
self.streams.capacity()
}
#[inline(always)]
pub fn is_empty(&self) -> bool {
self.streams.is_empty()
}
pub fn remove(&mut self, key: Key) -> bool {
let is_present = self.keys.remove(&key.0);
if is_present {
self.states[key.0].set_none();
self.streams.remove(key.0);
}
is_present
}
pub fn contains_key(&mut self, key: Key) -> bool {
self.keys.contains(&key.0)
}
pub fn reserve(&mut self, additional: usize) {
self.streams.reserve(additional);
let new_cap = self.streams.capacity();
self.wakers.resize(new_cap);
self.states.resize(new_cap);
}
}
impl<S: Stream> StreamGroup<S> {
pub fn insert(&mut self, stream: S) -> Key
where
S: Stream,
{
let index = self.streams.insert(stream);
self.keys.insert(index);
let new_cap = self.streams.capacity();
self.wakers.resize(new_cap);
self.states.resize(new_cap);
self.states[index].set_pending();
self.wakers.readiness().set_ready(index);
Key(index)
}
pub fn keyed(self) -> Keyed<S> {
Keyed { group: self }
}
}
impl<S: Stream> StreamGroup<S> {
fn poll_next_inner(
mut self: Pin<&mut Self>,
cx: &Context<'_>,
) -> Poll<Option<(Key, <S as Stream>::Item)>> {
let mut this = self.as_mut().project();
if this.streams.is_empty() {
return Poll::Ready(None);
}
let mut readiness = this.wakers.readiness();
readiness.set_waker(cx.waker());
if !readiness.any_ready() {
return Poll::Pending;
}
let mut ret = Poll::Pending;
let mut done_count = 0;
let stream_count = this.streams.len();
let states = this.states;
let streams = unsafe { this.streams.as_mut().get_unchecked_mut() };
for index in this.keys.iter().cloned() {
if states[index].is_pending() && readiness.clear_ready(index) {
#[allow(clippy::drop_non_drop)]
drop(readiness);
let mut cx = Context::from_waker(this.wakers.get(index).unwrap());
let stream = unsafe { Pin::new_unchecked(&mut streams[index]) };
match stream.poll_next(&mut cx) {
Poll::Ready(Some(item)) => {
ret = Poll::Ready(Some((Key(index), item)));
states[index] = PollState::Pending;
let mut readiness = this.wakers.readiness();
readiness.set_ready(index);
break;
}
Poll::Ready(None) => {
done_count += 1;
states[index] = PollState::None;
streams.remove(index);
this.key_removal_queue.push(index);
}
Poll::Pending => {}
};
readiness = this.wakers.readiness();
}
}
if !this.key_removal_queue.is_empty() {
for key in this.key_removal_queue.iter() {
this.keys.remove(key);
}
this.key_removal_queue.clear();
}
if done_count == stream_count {
ret = Poll::Ready(None);
}
ret
}
}
impl<S: Stream> Stream for StreamGroup<S> {
type Item = <S as Stream>::Item;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
match self.poll_next_inner(cx) {
Poll::Ready(Some((_key, item))) => Poll::Ready(Some(item)),
Poll::Ready(None) => Poll::Ready(None),
Poll::Pending => Poll::Pending,
}
}
}
impl<S: Stream> FromIterator<S> for StreamGroup<S> {
fn from_iter<T: IntoIterator<Item = S>>(iter: T) -> Self {
let iter = iter.into_iter();
let len = iter.size_hint().1.unwrap_or_default();
let mut this = Self::with_capacity(len);
for stream in iter {
this.insert(stream);
}
this
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub struct Key(usize);
#[derive(Debug)]
#[pin_project::pin_project]
pub struct Keyed<S: Stream> {
#[pin]
group: StreamGroup<S>,
}
impl<S: Stream> Deref for Keyed<S> {
type Target = StreamGroup<S>;
fn deref(&self) -> &Self::Target {
&self.group
}
}
impl<S: Stream> DerefMut for Keyed<S> {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.group
}
}
impl<S: Stream> Stream for Keyed<S> {
type Item = (Key, <S as Stream>::Item);
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let mut this = self.project();
this.group.as_mut().poll_next_inner(cx)
}
}
#[cfg(test)]
mod test {
use super::StreamGroup;
use futures_lite::{prelude::*, stream};
#[test]
fn smoke() {
futures_lite::future::block_on(async {
let mut group = StreamGroup::new();
group.insert(stream::once(2));
group.insert(stream::once(4));
let mut out = 0;
while let Some(num) = group.next().await {
out += num;
}
assert_eq!(out, 6);
assert_eq!(group.len(), 0);
assert!(group.is_empty());
});
}
#[test]
fn capacity_grow_on_insert() {
futures_lite::future::block_on(async {
let mut group = StreamGroup::new();
let cap = group.capacity();
group.insert(stream::once(1));
assert!(group.capacity() > cap);
});
}
}