use crate::{prelude::*, ConsumeLimitError};
#[derive(Copy, Clone, Hash, Ord, Eq, PartialEq, PartialOrd, Debug)]
pub struct Limit<C> {
inner: C,
remaining: usize,
}
impl<C> Limit<C> {
pub(crate) fn new(inner: C, limit: usize) -> Self {
Limit {
inner,
remaining: limit,
}
}
pub fn into_inner(self) -> C {
self.inner
}
pub fn remaining(&self) -> usize {
self.remaining
}
}
impl<C> AsRef<C> for Limit<C> {
fn as_ref(&self) -> &C {
&self.inner
}
}
impl<C: Consumer> Consumer for Limit<C> {
type Item = C::Item;
type Final = C::Final;
type Error = ConsumeLimitError<C::Error>;
async fn consume(&mut self, val: Either<Self::Item, Self::Final>) -> Result<(), Self::Error> {
match val {
Left(item) => match self.remaining.checked_sub(1) {
None => Result::Err(ConsumeLimitError::LimitReached),
Some(decremented) => {
self.remaining = decremented;
self.inner
.consume_item(item)
.await
.map_err(ConsumeLimitError::Inner)
}
},
Right(fin) => self
.inner
.consume_final(fin)
.await
.map_err(ConsumeLimitError::Inner),
}
}
async fn flush(&mut self) -> Result<(), Self::Error> {
self.inner.flush().await.map_err(ConsumeLimitError::Inner)
}
}
impl<C: BulkConsumer> BulkConsumer for Limit<C> {
async fn expose_slots_gracefully<F, R>(&mut self, f: F) -> Result<R, (F, Self::Error)>
where
F: AsyncFnOnce(&mut [Self::Item]) -> (usize, R),
{
if self.remaining == 0 {
return Err((f, ConsumeLimitError::LimitReached));
}
let mut f = Some(f);
let mut remaining = self.remaining;
let ret = self
.inner
.expose_slots(async |slots| {
let len = remaining.min(slots.len());
let inner_f = f.take().expect("This only runs once, so f is safe to take");
let (consumed, returned) = inner_f(&mut slots[..len]).await;
remaining -= consumed;
(consumed, returned)
})
.await;
self.remaining = remaining;
match ret {
Ok(r) => Ok(r),
Err(error) => Err((
f.take()
.expect("Wrapper closure wasn't called, so f is safe to take"),
ConsumeLimitError::Inner(error),
)),
}
}
}
#[cfg(test)]
mod tests {
use crate::prelude::*;
use crate::ConsumeLimitError;
#[test]
fn inner_error() {
let mut c = consumer::error_immediately("internal_error").to_limit(20);
pollster::block_on(async {
assert_eq!(
c.flush().await,
Err(ConsumeLimitError::Inner("internal_error"))
);
});
}
}