use std::{
fmt::Debug,
ops::{Add, AddAssign, Sub, SubAssign},
};
use crate::correlated_randomness::{
bundler::errors::BundlerError,
stream::{CorrelatedStreamError, PrefetchHandle},
};
pub mod errors;
pub type SizeOf<I> = <I as BundleIterator>::Size;
pub type ErrorOf<I> = <I as BundleIterator>::Error;
pub trait Bundler {
type Iterator: BundleIterator;
fn fetch(&mut self, size: &SizeOf<Self::Iterator>) -> Result<Self::Iterator, BundlerError>;
fn fetch_for<Consumer: BundleConsumer<Iterator = Self::Iterator>>(
&mut self,
consumer: &Consumer,
) -> Result<Self::Iterator, BundlerError> {
let size = consumer.required_preprocessing();
self.fetch(&size)
}
fn prefetch(&self, size: &SizeOf<Self::Iterator>) -> PrefetchHandle<ErrorOf<Self::Iterator>>;
fn prefetch_for<Consumer: BundleConsumer<Iterator = Self::Iterator>>(
&self,
consumer: &Consumer,
) -> PrefetchHandle<ErrorOf<Self::Iterator>> {
let size = consumer.required_preprocessing();
self.prefetch(&size)
}
}
impl<B: Bundler> Bundler for &mut B {
type Iterator = B::Iterator;
fn fetch(&mut self, size: &SizeOf<Self::Iterator>) -> Result<Self::Iterator, BundlerError> {
(**self).fetch(size)
}
fn prefetch(&self, size: &SizeOf<Self::Iterator>) -> PrefetchHandle<ErrorOf<Self::Iterator>> {
(**self).prefetch(size)
}
}
pub trait BundleConsumer {
type Iterator: BundleIterator;
fn required_preprocessing(&self) -> SizeOf<Self::Iterator>;
fn fetch_preprocessing_from<PB: Bundler<Iterator = Self::Iterator>>(
&self,
bundler: &mut PB,
) -> Result<Self::Iterator, BundlerError> {
let size = self.required_preprocessing();
bundler.fetch(&size)
}
fn prefetch_preprocessing_from<PB: Bundler<Iterator = Self::Iterator>>(
&self,
bundler: &PB,
) -> PrefetchHandle<ErrorOf<Self::Iterator>> {
let size = self.required_preprocessing();
bundler.prefetch(&size)
}
}
pub trait BundleIterator {
type Size: Debug + Clone + Default + Eq + Add + Sub + AddAssign + SubAssign;
type Error: Debug + Clone + From<CorrelatedStreamError> + Send + 'static;
fn len(&self) -> Self::Size;
fn is_empty(&self) -> bool {
self.len() == Self::Size::default()
}
}
#[cfg(test)]
mod tests {
use std::cell::Cell;
use super::*;
struct Bundle(usize);
impl BundleIterator for Bundle {
type Size = usize;
type Error = CorrelatedStreamError;
fn len(&self) -> usize {
self.0
}
}
#[derive(Default)]
struct Recorder(Cell<Option<usize>>);
impl Bundler for Recorder {
type Iterator = Bundle;
fn fetch(&mut self, size: &usize) -> Result<Bundle, BundlerError> {
Ok(Bundle(*size))
}
fn prefetch(&self, size: &usize) -> PrefetchHandle<CorrelatedStreamError> {
self.0.set(Some(*size));
PrefetchHandle::ready(Ok(()))
}
}
struct Consumer(usize);
impl BundleConsumer for Consumer {
type Iterator = Bundle;
fn required_preprocessing(&self) -> usize {
self.0
}
}
#[tokio::test]
async fn consumer_driven_prefetch_forwards_the_required_size() {
let bundler = Recorder::default();
bundler.prefetch_for(&Consumer(11)).await.unwrap();
assert_eq!(bundler.0.get(), Some(11));
let bundler = Recorder::default();
Consumer(23)
.prefetch_preprocessing_from(&bundler)
.await
.unwrap();
assert_eq!(bundler.0.get(), Some(23));
}
#[tokio::test]
async fn mut_ref_forwards_prefetch() {
let mut bundler = Recorder::default();
let forwarded = &mut bundler;
forwarded.prefetch(&5).await.unwrap();
assert_eq!(bundler.0.get(), Some(5));
}
}