use std::{borrow::BorrowMut, collections::HashMap, hash::Hash, time::Duration};
use futures::{Stream, StreamExt, stream};
use tokio::sync::{broadcast, mpsc};
#[allow(dead_code)]
#[allow(clippy::mutable_key_type)] pub fn get_item_counts<I>(items: I) -> HashMap<I::Item, usize>
where
I: IntoIterator,
I::Item: Hash + Eq,
{
items.into_iter().fold(HashMap::new(), |mut counts, item| {
let entry = counts.entry(item).or_insert(0);
*entry += 1;
counts
})
}
#[macro_export]
macro_rules! collect_stream {
($stream:expr, take=$take:expr, timeout=$timeout:expr $(,)?) => {{
use tokio::time;
let mut stream = &mut $stream;
let mut items = Vec::new();
loop {
if let Some(item) = time::timeout($timeout, futures::stream::StreamExt::next(stream))
.await
.expect(
format!(
"Timeout before stream could collect {} item(s). Got {} item(s).",
$take,
items.len()
)
.as_str(),
)
{
items.push(item);
if items.len() == $take {
break items;
}
} else {
break items;
}
}
}};
($stream:expr, timeout=$timeout:expr $(,)?) => {{
use tokio::time;
let mut stream = &mut $stream;
let mut items = Vec::new();
while let Some(item) = time::timeout($timeout, futures::stream::StreamExt::next($stream))
.await
.expect(format!("Timeout before stream was closed. Got {} items.", items.len()).as_str())
{
items.push(item);
}
items
}};
}
#[macro_export]
macro_rules! collect_recv {
($stream:expr, take=$take:expr, timeout=$timeout:expr $(,)?) => {{
use tokio::time;
let mut stream = &mut $stream;
let mut items = Vec::new();
loop {
let item = time::timeout($timeout, stream.recv()).await.expect(&format!(
"Timeout before stream could collect {} item(s). Got {} item(s).",
$take,
items.len()
));
items.push(item.expect(&format!("{}/{} recv ended early", items.len(), $take)));
if items.len() == $take {
break items;
}
}
}};
($stream:expr, timeout=$timeout:expr $(,)?) => {{
use tokio::time;
let mut stream = &mut $stream;
let mut items = Vec::new();
while let Some(item) = time::timeout($timeout, stream.recv())
.await
.expect(format!("Timeout before stream was closed. Got {} items.", items.len()).as_str())
{
items.push(item);
}
items
}};
}
#[macro_export]
macro_rules! collect_try_recv {
($stream:expr, take=$take:expr, timeout=$timeout:expr $(,)?) => {{
use tokio::time;
let mut stream = &mut $stream;
let mut items = Vec::new();
loop {
let item = time::timeout($timeout, stream.recv()).await.expect(&format!(
"Timeout before stream could collect {} item(s). Got {} item(s).",
$take,
items.len()
));
items.push(item.expect(&format!("{}/{} recv returned unexpected result", items.len(), $take)));
if items.len() == $take {
break items;
}
}
}};
($stream:expr, timeout=$timeout:expr $(,)?) => {{
use tokio::time;
let mut stream = &mut $stream;
let mut items = Vec::new();
while let Ok(item) = time::timeout($timeout, stream.recv())
.await
.expect(format!("Timeout before stream was closed. Got {} items.", items.len()).as_str())
{
items.push(item);
}
items
}};
}
#[macro_export]
macro_rules! collect_stream_count {
($stream:expr, take=$take:expr, timeout=$timeout:expr$(,)?) => {{
use std::collections::HashMap;
let items = $crate::collect_stream!($stream, take = $take, timeout = $timeout);
$crate::streams::get_item_counts(items)
}};
($stream:expr, timeout=$timeout:expr $(,)?) => {{
use std::collections::HashMap;
let items = $crate::collect_stream!($stream, timeout = $timeout);
$crate::streams::get_item_counts(items)
}};
}
pub async fn assert_in_stream<S, P, R>(stream: &mut S, mut predicate: P, timeout: Duration) -> R
where
S: Stream + Unpin,
P: FnMut(S::Item) -> Option<R>,
{
loop {
if let Some(item) = tokio::time::timeout(timeout, stream.next())
.await
.expect("Timeout before stream emitted")
{
if let Some(r) = (predicate)(item) {
break r;
}
} else {
panic!("Predicate did not return true before the stream ended");
}
}
}
pub async fn assert_in_mpsc<T, P, R>(rx: &mut mpsc::Receiver<T>, mut predicate: P, timeout: Duration) -> R
where P: FnMut(T) -> Option<R> {
loop {
if let Some(item) = tokio::time::timeout(timeout, rx.recv())
.await
.expect("Timeout before stream emitted")
{
if let Some(r) = (predicate)(item) {
break r;
}
} else {
panic!("Predicate did not return true before the mpsc stream ended");
}
}
}
pub async fn assert_in_broadcast<T, P, R>(rx: &mut broadcast::Receiver<T>, mut predicate: P, timeout: Duration) -> R
where
P: FnMut(T) -> Option<R>,
T: Clone,
{
loop {
if let Ok(item) = tokio::time::timeout(timeout, rx.recv())
.await
.expect("Timeout before stream emitted")
{
if let Some(r) = (predicate)(item) {
break r;
}
} else {
panic!("Predicate did not return true before the broadcast channel ended");
}
}
}
pub fn convert_mpsc_to_stream<T>(rx: &mut mpsc::Receiver<T>) -> impl Stream<Item = T> + '_ {
stream::unfold(rx, |rx| async move { rx.recv().await.map(|t| (t, rx)) })
}
pub fn convert_unbounded_mpsc_to_stream<T>(rx: &mut mpsc::UnboundedReceiver<T>) -> impl Stream<Item = T> + '_ {
stream::unfold(rx, |rx| async move { rx.recv().await.map(|t| (t, rx)) })
}
pub fn convert_broadcast_to_stream<'a, T, S>(rx: S) -> impl Stream<Item = Result<T, broadcast::error::RecvError>> + 'a
where
T: Clone + Send + 'static,
S: BorrowMut<broadcast::Receiver<T>> + 'a,
{
stream::unfold(rx, |mut rx| async move {
Some(rx.borrow_mut().recv().await).map(|t| (t, rx))
})
}