use std::num::NonZeroUsize;
use std::pin::Pin;
use futures::{Stream, StreamExt, stream};
pub fn chunk_by_bounded<S, K, F>(
stream: S,
max: NonZeroUsize,
key: F,
) -> impl Stream<Item = (K, Vec<S::Item>)>
where
S: Stream + Unpin,
K: PartialEq,
F: FnMut(&S::Item) -> K,
{
let max = max.get();
stream::unfold((stream.peekable(), key), move |(mut stream, mut key)| async move {
let target = key(Pin::new(&mut stream).peek().await?);
let mut chunk = Vec::new();
while chunk.len() < max {
match Pin::new(&mut stream).peek().await {
Some(item) if key(item) == target => {}
_ => break,
}
chunk.push(stream.next().await.expect("peeked item is present"));
}
Some(((target, chunk), (stream, key)))
})
}
#[cfg(test)]
mod tests {
use std::num::NonZeroUsize;
use futures::executor::block_on;
use futures::{StreamExt, stream};
use proptest::prelude::*;
use rstest::rstest;
use super::chunk_by_bounded;
fn chunks_of(items: Vec<i32>, max: usize) -> Vec<Vec<i32>> {
let max = NonZeroUsize::new(max).expect("test max is non-zero");
block_on(
chunk_by_bounded(stream::iter(items), max, |x| *x).map(|(_, chunk)| chunk).collect(),
)
}
#[rstest]
#[case::empty(vec![], 3, vec![])]
#[case::single(vec![5], 3, vec![vec![5]])]
#[case::max_one_splits_every_item(vec![1, 1, 2], 1, vec![vec![1], vec![1], vec![2]])]
#[case::alternating_keys(vec![1, 2, 1], 5, vec![vec![1], vec![2], vec![1]])]
#[case::run_exactly_max(vec![1, 1], 2, vec![vec![1, 1]])]
#[case::run_longer_than_max(vec![1, 1, 1, 1, 1], 2, vec![vec![1, 1], vec![1, 1], vec![1]])]
#[case::mixed(vec![1, 1, 1, 2, 2, 3], 5, vec![vec![1, 1, 1], vec![2, 2], vec![3]])]
fn chunks_match_expected(
#[case] items: Vec<i32>,
#[case] max: usize,
#[case] expected: Vec<Vec<i32>>,
) {
assert_eq!(chunks_of(items, max), expected);
}
#[test]
fn groups_by_key_not_value() {
let chunks: Vec<(i32, Vec<i32>)> = block_on(
chunk_by_bounded(stream::iter([2, 4, 3, 6, 5]), NonZeroUsize::new(5).unwrap(), |x| {
x % 2
})
.collect(),
);
assert_eq!(chunks, vec![(0, vec![2, 4]), (1, vec![3]), (0, vec![6]), (1, vec![5])]);
}
proptest! {
#[test]
fn holds_invariants(items in prop::collection::vec(0i32..4, 0..50), max in 1usize..8) {
let chunks = chunks_of(items.clone(), max);
let flat: Vec<i32> = chunks.iter().flatten().copied().collect();
prop_assert_eq!(&flat, &items);
for chunk in &chunks {
prop_assert!(!chunk.is_empty());
prop_assert!(chunk.len() <= max);
prop_assert!(chunk.iter().all(|x| x == &chunk[0]));
}
for pair in chunks.windows(2) {
if pair[0].len() < max {
prop_assert_ne!(pair[0].last().unwrap(), pair[1].first().unwrap());
}
}
}
}
}