use std::pin::Pin;
use futures::future;
use futures::future::BoxFuture;
use futures::ready;
use futures::stream;
use futures::task::Context;
use futures::task::Poll;
use futures::Future;
use futures::FutureExt;
use futures::Stream;
use futures::StreamExt;
use futures::TryStream;
use pin_project::pin_project;
#[derive(Clone, Copy, Debug)]
pub struct BufferedParams {
pub weight_limit: u64,
pub buffer_size: usize,
}
#[pin_project]
pub struct WeightLimitedBufferedStream<'a, S, I> {
#[pin]
queue: stream::FuturesOrdered<BoxFuture<'a, (I, u64)>>,
current_weight: u64,
weight_limit: u64,
max_buffer_size: usize,
#[pin]
stream: stream::Fuse<S>,
}
impl<S, I> WeightLimitedBufferedStream<'_, S, I>
where
S: Stream,
{
pub fn new(params: BufferedParams, stream: S) -> Self {
Self {
queue: stream::FuturesOrdered::new(),
current_weight: 0,
weight_limit: params.weight_limit,
max_buffer_size: params.buffer_size,
stream: stream.fuse(),
}
}
}
impl<'a, S, Fut, I: 'a> Stream for WeightLimitedBufferedStream<'a, S, I>
where
S: Stream<Item = (Fut, u64)>,
Fut: Future<Output = I> + Send + 'a,
{
type Item = I;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let mut this = self.project();
while this.queue.len() < *this.max_buffer_size && this.current_weight < this.weight_limit {
let future = match this.stream.as_mut().poll_next(cx) {
Poll::Ready(Some((f, weight))) => {
*this.current_weight += weight;
f.map(move |val| (val, weight)).boxed()
}
Poll::Ready(None) | Poll::Pending => break,
};
this.queue.push_back(future);
}
if let Some((val, weight)) = ready!(this.queue.poll_next(cx)) {
*this.current_weight -= weight;
return Poll::Ready(Some(val));
}
if this.stream.is_done() {
Poll::Ready(None)
} else {
Poll::Pending
}
}
}
#[pin_project]
pub struct WeightLimitedBufferedTryStream<'a, S, I, E> {
#[pin]
queue: stream::FuturesOrdered<BoxFuture<'a, (Result<I, E>, u64)>>,
current_weight: u64,
weight_limit: u64,
max_buffer_size: usize,
#[pin]
stream: stream::Fuse<S>,
}
impl<S, I, E> WeightLimitedBufferedTryStream<'_, S, I, E>
where
S: TryStream,
{
pub fn new(params: BufferedParams, stream: S) -> Self {
Self {
queue: stream::FuturesOrdered::new(),
current_weight: 0,
weight_limit: params.weight_limit,
max_buffer_size: params.buffer_size,
stream: stream.fuse(),
}
}
}
impl<'a, S, Fut, I: 'a, E> Stream for WeightLimitedBufferedTryStream<'a, S, I, E>
where
S: Stream<Item = Result<(Fut, u64), E>>,
Fut: Future<Output = Result<I, E>> + Send + 'a,
E: Send + 'a,
I: Send,
{
type Item = Result<I, E>;
fn poll_next(self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
let mut this = self.project();
while this.queue.len() < *this.max_buffer_size && this.current_weight < this.weight_limit {
let future = match this.stream.as_mut().poll_next(cx) {
Poll::Ready(Some(Ok((f, weight)))) => {
*this.current_weight += weight;
f.map(move |val| (val, weight)).boxed()
}
Poll::Ready(Some(Err(e))) => {
future::ready((Err(e), 0u64)).boxed()
}
Poll::Ready(None) | Poll::Pending => break,
};
this.queue.push_back(future);
}
if let Some((val, weight)) = ready!(this.queue.poll_next(cx)) {
*this.current_weight -= weight;
return Poll::Ready(Some(val));
}
if this.stream.is_done() {
Poll::Ready(None)
} else {
Poll::Pending
}
}
}
#[cfg(test)]
mod test {
use std::sync::atomic::AtomicUsize;
use std::sync::atomic::Ordering;
use std::sync::Arc;
use futures::future;
use futures::future::BoxFuture;
use futures::stream;
use futures::stream::BoxStream;
use futures::FutureExt;
use futures::StreamExt;
use super::*;
type TestStream = BoxStream<'static, (BoxFuture<'static, ()>, u64)>;
fn create_stream() -> (Arc<AtomicUsize>, TestStream) {
let s: TestStream = stream::iter(vec![
(future::ready(()).boxed(), 100),
(future::ready(()).boxed(), 2),
(future::ready(()).boxed(), 7),
])
.boxed();
let counter = Arc::new(AtomicUsize::new(0));
(
counter.clone(),
s.inspect({
move |_val| {
counter.fetch_add(1, Ordering::SeqCst);
}
})
.boxed(),
)
}
#[tokio::test]
async fn test_too_much_weight_to_do_in_one_go() {
let (counter, s) = create_stream();
let params = BufferedParams {
weight_limit: 10,
buffer_size: 10,
};
let s = WeightLimitedBufferedStream::new(params, s);
if let (Some(()), s) = s.into_future().await {
assert_eq!(counter.load(Ordering::SeqCst), 1);
assert_eq!(s.collect::<Vec<()>>().await.len(), 2);
assert_eq!(counter.load(Ordering::SeqCst), 3);
} else {
panic!("Stream did not produce even a single value");
}
}
#[tokio::test]
async fn test_all_in_one_go() {
let (counter, s) = create_stream();
let params = BufferedParams {
weight_limit: 200,
buffer_size: 10,
};
let s = WeightLimitedBufferedStream::new(params, s);
if let (Some(()), s) = s.into_future().await {
assert_eq!(counter.load(Ordering::SeqCst), 3);
assert_eq!(s.collect::<Vec<()>>().await.len(), 2);
assert_eq!(counter.load(Ordering::SeqCst), 3);
} else {
panic!("Stream did not produce even a single value");
}
}
#[tokio::test]
async fn test_too_much_items_to_do_in_one_go() {
let (counter, s) = create_stream();
let params = BufferedParams {
weight_limit: 1000,
buffer_size: 2,
};
let s = WeightLimitedBufferedStream::new(params, s);
if let (Some(()), s) = s.into_future().await {
assert_eq!(counter.load(Ordering::SeqCst), 2);
assert_eq!(s.collect::<Vec<()>>().await.len(), 2);
assert_eq!(counter.load(Ordering::SeqCst), 3);
} else {
panic!("Stream did not produce even a single value");
}
}
type Error = String;
type TestTryStream =
BoxStream<'static, Result<(BoxFuture<'static, Result<(), Error>>, u64), Error>>;
fn counted_try_stream(s: TestTryStream) -> (Arc<AtomicUsize>, TestTryStream) {
let counter = Arc::new(AtomicUsize::new(0));
(
counter.clone(),
s.inspect({
move |_val| {
counter.fetch_add(1, Ordering::SeqCst);
}
})
.boxed(),
)
}
fn create_try_stream_all_good() -> (Arc<AtomicUsize>, TestTryStream) {
let s: TestTryStream = stream::iter(vec![
Ok((future::ready(Ok(())).boxed(), 100)),
Ok((future::ready(Ok(())).boxed(), 2)),
Ok((future::ready(Ok(())).boxed(), 7)),
])
.boxed();
counted_try_stream(s)
}
#[tokio::test]
async fn test_try_all_in_one_go() {
let (counter, s) = create_try_stream_all_good();
let params = BufferedParams {
weight_limit: 200,
buffer_size: 10,
};
let s = WeightLimitedBufferedTryStream::new(params, s);
if let (Some(Ok(())), s) = s.into_future().await {
assert_eq!(counter.load(Ordering::SeqCst), 3);
assert_eq!(s.collect::<Vec<_>>().await.len(), 2);
assert_eq!(counter.load(Ordering::SeqCst), 3);
} else {
panic!("Stream did not produce even a single value");
}
}
#[tokio::test]
async fn test_try_too_much_weight_to_do_in_one_go() {
let (counter, s) = create_try_stream_all_good();
let params = BufferedParams {
weight_limit: 10,
buffer_size: 10,
};
let s = WeightLimitedBufferedTryStream::new(params, s);
if let (Some(Ok(())), s) = s.into_future().await {
assert_eq!(counter.load(Ordering::SeqCst), 1);
assert_eq!(s.collect::<Vec<_>>().await.len(), 2);
assert_eq!(counter.load(Ordering::SeqCst), 3);
} else {
panic!("Stream did not produce even a single value");
}
}
#[tokio::test]
async fn test_try_too_much_items_to_do_in_one_go() {
let (counter, s) = create_try_stream_all_good();
let params = BufferedParams {
weight_limit: 1000,
buffer_size: 2,
};
let s = WeightLimitedBufferedTryStream::new(params, s);
if let (Some(Ok(())), s) = s.into_future().await {
assert_eq!(counter.load(Ordering::SeqCst), 2);
assert_eq!(s.collect::<Vec<_>>().await.len(), 2);
assert_eq!(counter.load(Ordering::SeqCst), 3);
} else {
panic!("Stream did not produce even a single value");
}
}
fn create_try_stream_fail_external() -> (Arc<AtomicUsize>, TestTryStream) {
let s: TestTryStream = stream::iter(vec![
Ok((future::ready(Ok(())).boxed(), 100)),
Err("failed to calculate weight".to_string()),
Ok((future::ready(Ok(())).boxed(), 7)),
])
.boxed();
counted_try_stream(s)
}
#[tokio::test]
async fn test_try_fail_to_calculate_weight() {
let (counter, s) = create_try_stream_fail_external();
let params = BufferedParams {
weight_limit: 1000,
buffer_size: 2,
};
let s = WeightLimitedBufferedTryStream::new(params, s);
if let (Some(Ok(())), s) = s.into_future().await {
assert_eq!(counter.load(Ordering::SeqCst), 2);
let v = s.collect::<Vec<Result<_, _>>>().await;
assert!(v[0].is_err());
assert!(
v[0].clone()
.unwrap_err()
.contains("failed to calculate weight")
);
assert_eq!(v[1], Ok(()));
assert_eq!(v.len(), 2);
assert_eq!(counter.load(Ordering::SeqCst), 3);
} else {
panic!("Stream did not produce even a single value");
}
}
fn create_try_stream_fail_internal() -> (Arc<AtomicUsize>, TestTryStream) {
let s: TestTryStream = stream::iter(vec![
Ok((future::ready(Ok(())).boxed(), 100)),
Ok((
future::ready(Err("failed to produce interesting value".to_string())).boxed(),
2,
)),
Ok((future::ready(Ok(())).boxed(), 7)),
])
.boxed();
counted_try_stream(s)
}
#[tokio::test]
async fn test_try_fail_to_calculate_inner_value() {
let (counter, s) = create_try_stream_fail_internal();
let params = BufferedParams {
weight_limit: 1000,
buffer_size: 2,
};
let s = WeightLimitedBufferedTryStream::new(params, s);
if let (Some(Ok(())), s) = s.into_future().await {
assert_eq!(counter.load(Ordering::SeqCst), 2);
let v = s.collect::<Vec<Result<_, _>>>().await;
assert!(v[0].is_err());
assert!(
v[0].clone()
.unwrap_err()
.contains("failed to produce interesting value")
);
assert_eq!(v[1], Ok(()));
assert_eq!(v.len(), 2);
assert_eq!(counter.load(Ordering::SeqCst), 3);
} else {
panic!("Stream did not produce even a single value");
}
}
}