use crate::{prelude::*, ProduceLimitError};
#[derive(Copy, Clone, Hash, Ord, Eq, PartialEq, PartialOrd, Debug)]
pub struct Limit<P> {
inner: P,
remaining: usize,
}
impl<P> Limit<P> {
pub(crate) fn new(inner: P, limit: usize) -> Self {
Limit {
inner,
remaining: limit,
}
}
pub fn into_inner(self) -> P {
self.inner
}
pub fn remaining(&self) -> usize {
self.remaining
}
}
impl<P> AsRef<P> for Limit<P> {
fn as_ref(&self) -> &P {
&self.inner
}
}
impl<P: Producer> Producer for Limit<P> {
type Item = P::Item;
type Final = P::Final;
type Error = ProduceLimitError<P::Error>;
async fn produce(&mut self) -> Result<Either<Self::Item, Self::Final>, Self::Error> {
match self.remaining.checked_sub(1) {
None => Result::Err(ProduceLimitError::LimitReached),
Some(decremented) => {
self.remaining = decremented;
self.inner
.produce()
.await
.map_err(|err| ProduceLimitError::Inner(err))
}
}
}
async fn slurp(&mut self) -> Result<(), Self::Error> {
self.inner
.slurp()
.await
.map_err(|err| ProduceLimitError::Inner(err))
}
}
impl<P: BulkProducer> BulkProducer for Limit<P> {
async fn expose_items_gracefully<F, R>(
&mut self,
f: F,
) -> Result<Either<R, (F, Self::Final)>, (F, Self::Error)>
where
F: AsyncFnOnce(&[Self::Item]) -> (usize, R),
{
if self.remaining == 0 {
return Err((f, ProduceLimitError::LimitReached));
}
let mut f = Some(f);
let mut remaining = self.remaining;
let ret = self
.inner
.expose_items(async |items| {
let len = remaining.min(items.len());
remaining = remaining - len;
let inner_f = f
.take()
.expect("This should only run once, so f should be safe to take");
inner_f(&items[..len]).await
})
.await;
self.remaining = remaining;
match ret {
Ok(Left(r)) => Ok(Left(r)),
Ok(Right(fin)) => Ok(Right((
f.take()
.expect("Wrapper closure wasn't called, so f is safe to take"),
fin,
))),
Err(error) => Err((
f.take()
.expect("Wrapper closure wasn't called, so f is safe to take"),
ProduceLimitError::Inner(error),
)),
}
}
}
#[cfg(test)]
mod tests {
use crate::prelude::*;
use crate::ProduceLimitError;
#[test]
fn inner_error() {
let mut p = producer::error_immediately("internal_error").to_limit(20);
pollster::block_on(async {
assert_eq!(
p.produce().await,
Err(ProduceLimitError::Inner("internal_error"))
);
});
}
}