pub mod config;
pub mod messages;
use std::{
cell::Cell,
collections::VecDeque,
marker::PhantomData,
sync::{
atomic::{AtomicU64, Ordering},
Arc,
},
};
pub use config::{Buffer, BufferConfig, SharedBufferConfig};
use log::{debug, error, info};
pub use messages::PrefetchHandle;
use parking_lot::RwLock;
use tokio::{
sync::{
mpsc::{self, UnboundedReceiver, UnboundedSender},
oneshot,
},
task::JoinHandle,
};
use crate::correlated_randomness::{
generator::CorrelationGenerator,
stream::{
buffered::{config::try_read_config, messages::Command},
errors::CorrelatedStreamError,
futures::Next,
CorrelatedStream,
NextVec,
ResyncHandle,
},
CorrelatedBatch,
};
enum Work {
Generate(usize),
Skip(usize),
}
pub struct BufferedStream<PB: CorrelatedBatch, E> {
command_sender: UnboundedSender<Command<PB, E>>,
_unsync_marker: PhantomData<Cell<()>>,
config: Arc<RwLock<BufferConfig>>,
position: Arc<AtomicU64>,
buffered: Arc<AtomicU64>,
dispatcher_handle: JoinHandle<()>,
generator_handle: JoinHandle<()>,
}
impl<PB: CorrelatedBatch, E> Buffer for BufferedStream<PB, E> {
fn config(&self) -> &Arc<RwLock<BufferConfig>> {
&self.config
}
}
impl<
PB: CorrelatedBatch,
E: From<CorrelatedStreamError> + Clone + Send + std::fmt::Debug + 'static,
> BufferedStream<PB, E>
{
pub fn new<G: CorrelationGenerator<PB> + Send + 'static>(
generator: G,
net: G::Net,
config: BufferConfig,
) -> Self
where
E: From<G::Error>,
{
Self::new_with_shared_config(generator, net, Arc::new(RwLock::new(config)))
}
pub fn new_on_demand<G: CorrelationGenerator<PB> + Send + 'static>(
generator: G,
net: G::Net,
) -> Self
where
E: From<G::Error>,
{
Self::new(
generator,
net,
BufferConfig::lazy(BufferConfig::UNBOUNDED, 0),
)
}
pub fn new_with_shared_config<G: CorrelationGenerator<PB> + Send + 'static>(
generator: G,
net: G::Net,
config: Arc<RwLock<BufferConfig>>,
) -> Self
where
E: From<G::Error>,
{
let (cmd_tx, cmd_rx) = mpsc::unbounded_channel::<Command<PB, E>>();
let (work_tx, work_rx) = mpsc::unbounded_channel::<Work>();
let (items_tx, items_rx) = mpsc::unbounded_channel::<Result<Vec<PB::Item>, E>>();
let (skip_tx, skip_rx) = mpsc::unbounded_channel::<Result<(), E>>();
let position = Arc::new(AtomicU64::new(0));
let buffered = Arc::new(AtomicU64::new(0));
let generator_handle =
tokio::spawn(generator_loop(generator, net, work_rx, items_tx, skip_tx));
let dispatcher_handle = tokio::spawn(dispatcher_loop(
cmd_rx,
work_tx,
items_rx,
skip_rx,
config.clone(),
position.clone(),
buffered.clone(),
G::SUPPORTS_UNILATERAL_SKIP,
));
Self {
command_sender: cmd_tx,
_unsync_marker: PhantomData,
config,
position,
buffered,
dispatcher_handle,
generator_handle,
}
}
pub async fn stop(self) {
let Self {
command_sender,
dispatcher_handle,
generator_handle,
..
} = self;
drop(command_sender);
let _ = dispatcher_handle.await;
let _ = generator_handle.await;
}
}
async fn generator_loop<
PB: CorrelatedBatch,
G: CorrelationGenerator<PB> + Send,
E: From<G::Error> + Send,
>(
mut generator: G,
mut net: G::Net,
mut work_rx: UnboundedReceiver<Work>,
items_tx: UnboundedSender<Result<Vec<PB::Item>, E>>,
skip_tx: UnboundedSender<Result<(), E>>,
) {
let log_prefix = format!("<Generator<{}>>", std::any::type_name::<PB>());
while let Some(work) = work_rx.recv().await {
match work {
Work::Generate(n) => {
debug!("{log_prefix} generating {n} elements");
let result = generator.run_for(n, &mut net).await.map_err(E::from);
let stop = result.is_err();
if items_tx.send(result).is_err() || stop {
break;
}
}
Work::Skip(n) => {
debug!("{log_prefix} skipping {n} elements");
let result = generator.skip(n, &mut net).await.map_err(E::from);
let stop = result.is_err();
if skip_tx.send(result).is_err() || stop {
break;
}
}
}
}
info!("{log_prefix} exiting");
}
type ItemsCollected<PB> = Vec<<PB as IntoIterator>::Item>;
type TotalNeeded = usize;
type BatchSender<PB, E> = oneshot::Sender<Result<Vec<<PB as IntoIterator>::Item>, E>>;
fn maybe_skip(
pending_skip: &mut usize,
work_in_flight: &mut bool,
work_tx: &UnboundedSender<Work>,
log_prefix: &str,
) {
if *pending_skip > 0 && !*work_in_flight {
debug!("{log_prefix} requesting skip of {} elements", *pending_skip);
*work_in_flight = true;
let _ = work_tx.send(Work::Skip(*pending_skip));
*pending_skip = 0;
}
}
#[allow(clippy::too_many_arguments)]
fn maybe_generate<PB: CorrelatedBatch, E: From<CorrelatedStreamError>>(
work_in_flight: &mut bool,
pending_resync: &Option<(usize, oneshot::Sender<Result<(), E>>)>,
pending_skip: usize,
config: &Arc<RwLock<BufferConfig>>,
pending_batches: &VecDeque<(ItemsCollected<PB>, TotalNeeded, BatchSender<PB, E>)>,
buffer_len: usize,
prefetch_demand: usize,
work_tx: &UnboundedSender<Work>,
log_prefix: &str,
) -> Result<(), E> {
if *work_in_flight || pending_resync.is_some() || pending_skip != 0 {
return Ok(());
}
let batch_shortfall: usize = pending_batches
.iter()
.map(|(collected, needed, _)| needed.saturating_sub(collected.len() + buffer_len))
.sum();
let buf_after = buffer_len.saturating_sub(batch_shortfall);
let need = batch_shortfall
+ try_read_config(config)?
.refill_threshold()
.saturating_sub(buf_after)
.max(prefetch_demand);
if need > 0 {
debug!("{log_prefix} requesting generation of {need} items");
*work_in_flight = true;
let _ = work_tx.send(Work::Generate(need));
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
async fn dispatcher_loop<
PB: CorrelatedBatch,
E: From<CorrelatedStreamError> + Clone + Send + std::fmt::Debug,
>(
mut cmd_rx: UnboundedReceiver<Command<PB, E>>,
work_tx: UnboundedSender<Work>,
mut items_rx: UnboundedReceiver<Result<Vec<PB::Item>, E>>,
mut skip_rx: UnboundedReceiver<Result<(), E>>,
config: Arc<RwLock<BufferConfig>>,
shared_position: Arc<AtomicU64>,
shared_buffered: Arc<AtomicU64>,
supports_skip: bool,
) {
let log_prefix = format!("<Dispatcher<{}>>", std::any::type_name::<PB>());
let initial_cap = try_read_config(&config)
.map(|c| c.capacity())
.unwrap_or(0)
.min(1 << 12);
let mut buffer: VecDeque<PB::Item> = VecDeque::with_capacity(initial_cap);
let mut prefetch_demand: usize = 0;
let mut prefetch_completions: VecDeque<(usize, oneshot::Sender<Result<(), E>>)> =
VecDeque::new();
let mut work_in_flight = false;
let mut pending_batches: VecDeque<(ItemsCollected<PB>, TotalNeeded, BatchSender<PB, E>)> =
VecDeque::new();
let mut pending_resync: Option<(usize, oneshot::Sender<Result<(), E>>)> = None;
let mut pending_skip: usize = 0;
let mut shutdown_err: Option<E> = None;
macro_rules! lock_cfg {
() => {
match try_read_config(&config) {
Ok(g) => g,
Err(e) => {
error!("{log_prefix} config lock timeout, shutting down");
shutdown_err = Some(e.into());
break;
}
}
};
}
macro_rules! outstanding_demand {
() => {{
let batch_shortfall: usize = pending_batches
.iter()
.map(|(collected, needed, _)| needed.saturating_sub(collected.len()))
.sum();
(batch_shortfall + prefetch_demand).saturating_sub(buffer.len())
}};
}
let set_buffered =
|buffer_len: usize| shared_buffered.store(buffer_len as u64, Ordering::Release);
loop {
tokio::select! {
cmd = cmd_rx.recv() => {
let Some(cmd) = cmd else {
info!("{log_prefix} command channel closed, shutting down");
break;
};
match cmd {
Command::RequestN { n_elements, completion } => {
debug!("{log_prefix} batch request for {n_elements} items");
let cap = lock_cfg!().capacity();
if outstanding_demand!() + n_elements > cap {
let _ = completion.send(Err(CorrelatedStreamError::RateLimitExceeded.into()));
} else if buffer.len() >= n_elements {
let items: Vec<_> = buffer.drain(..n_elements).collect();
shared_position.fetch_add(n_elements as u64, Ordering::Release);
set_buffered(buffer.len());
let _ = completion.send(Ok(items));
} else {
let collected: Vec<_> = buffer.drain(..).collect();
shared_position.fetch_add(n_elements as u64, Ordering::Release);
set_buffered(buffer.len());
pending_batches.push_back((collected, n_elements, completion));
}
}
Command::Resync { target, completion } => {
let position = shared_position.load(Ordering::Relaxed);
debug!("{log_prefix} resync to {target} (position {position})");
if pending_resync.is_some() {
let _ = completion.send(Err(CorrelatedStreamError::ResyncInProgress.into()));
} else if target < position {
let _ = completion.send(Err(CorrelatedStreamError::ResyncRewind {
current: position,
target,
}
.into()));
} else {
let skip = (target - position) as usize;
let drained = skip.min(buffer.len());
buffer.drain(..drained);
shared_position.fetch_add(drained as u64, Ordering::Release);
set_buffered(buffer.len());
let deficit = skip - drained;
if deficit == 0 {
let _ = completion.send(Ok(()));
} else if !supports_skip {
let _ = completion.send(Err(CorrelatedStreamError::ResyncUnsupported {
generated: position + drained as u64,
target,
}
.into()));
} else {
pending_resync = Some((deficit, completion));
pending_skip = deficit;
maybe_skip(&mut pending_skip, &mut work_in_flight, &work_tx, &log_prefix);
}
}
}
Command::Prefetch { n_elements, completion } => {
debug!("{log_prefix} prefetch {n_elements} items");
if buffer.len() >= n_elements {
let _ = completion.send(Ok(()));
} else {
let cap = lock_cfg!().capacity();
if outstanding_demand!() + n_elements > cap {
let _ = completion.send(Err(CorrelatedStreamError::RateLimitExceeded.into()));
} else {
let deficit = n_elements - buffer.len();
prefetch_demand += deficit;
prefetch_completions.push_back((deficit, completion));
}
}
}
}
}
result = items_rx.recv() => {
let Some(result) = result else {
info!("{log_prefix} generator channel closed, shutting down");
break;
};
match result {
Ok(items) => {
work_in_flight = false;
let generated = items.len();
debug!("{log_prefix} received {generated} items (pending_batches: {}, buffer: {})",
pending_batches.iter().map(|(c, n, _)| n - c.len()).sum::<usize>(),
buffer.len());
let mut iter = items.into_iter();
while let Some((ref mut collected, needed, _)) = pending_batches.front_mut() {
let shortfall = *needed - collected.len();
collected.extend(iter.by_ref().take(shortfall));
if collected.len() < *needed { break; } let (collected, _, tx) = pending_batches.pop_front().unwrap();
let _ = tx.send(Ok(collected));
}
buffer.extend(iter);
set_buffered(buffer.len());
prefetch_demand = prefetch_demand.saturating_sub(generated);
let mut credit = generated;
while credit > 0 {
let Some((deficit, _)) = prefetch_completions.front_mut() else { break; };
let used = (*deficit).min(credit);
*deficit -= used;
credit -= used;
if *deficit > 0 { break; }
let (_, tx) = prefetch_completions.pop_front().unwrap();
let _ = tx.send(Ok(()));
}
maybe_skip(&mut pending_skip, &mut work_in_flight, &work_tx, &log_prefix);
}
Err(e) => {
error!("{log_prefix} generation error, shutting down: {e:?}");
for (_, _, tx) in pending_batches.drain(..) { let _ = tx.send(Err(e.clone())); }
for (_, tx) in prefetch_completions.drain(..) { let _ = tx.send(Err(e.clone())); }
if let Some((_, tx)) = pending_resync.take() { let _ = tx.send(Err(e.clone())); }
return;
}
}
}
result = skip_rx.recv() => {
let Some(result) = result else {
info!("{log_prefix} skip channel closed, shutting down");
break;
};
work_in_flight = false;
match result {
Ok(()) => {
if let Some((deficit, tx)) = pending_resync.take() {
shared_position.fetch_add(deficit as u64, Ordering::Release);
let _ = tx.send(Ok(()));
}
}
Err(e) => {
error!("{log_prefix} skip error, shutting down: {e:?}");
for (_, _, tx) in pending_batches.drain(..) { let _ = tx.send(Err(e.clone())); }
for (_, tx) in prefetch_completions.drain(..) { let _ = tx.send(Err(e.clone())); }
if let Some((_, tx)) = pending_resync.take() { let _ = tx.send(Err(e.clone())); }
return;
}
}
}
}
if let Err(e) = maybe_generate::<PB, E>(
&mut work_in_flight,
&pending_resync,
pending_skip,
&config,
&pending_batches,
buffer.len(),
prefetch_demand,
&work_tx,
&log_prefix,
) {
error!("{log_prefix} config lock timeout, shutting down");
shutdown_err = Some(e);
break;
}
}
let final_err: E = shutdown_err.unwrap_or_else(|| CorrelatedStreamError::StreamClosed.into());
for (_, _, tx) in pending_batches.drain(..) {
let _ = tx.send(Err(final_err.clone()));
}
for (_, tx) in prefetch_completions.drain(..) {
let _ = tx.send(Err(final_err.clone()));
}
if let Some((_, tx)) = pending_resync.take() {
let _ = tx.send(Err(final_err.clone()));
}
}
impl<
PB: CorrelatedBatch,
E: From<CorrelatedStreamError> + Clone + Send + std::fmt::Debug + 'static,
> CorrelatedStream<PB::Item> for BufferedStream<PB, E>
{
type Error = E;
fn next_n(&self, n_elements: usize) -> Result<NextVec<PB::Item, E>, CorrelatedStreamError> {
if n_elements == 0 {
return Ok(NextVec::default());
}
let max_allowed = self.max_request_size()?;
if n_elements > max_allowed {
return Err(CorrelatedStreamError::RequestTooLarge {
requested: n_elements,
max_allowed,
});
}
let (tx, rx) = oneshot::channel();
self.command_sender
.send(Command::RequestN {
n_elements,
completion: tx,
})
.map_err(|e| CorrelatedStreamError::SendError(e.to_string()))?;
Ok(NextVec {
future: Next(rx),
size: n_elements,
})
}
fn prefetch_n(&self, n_elements: usize) -> PrefetchHandle<E> {
let (tx, rx) = oneshot::channel();
let max = match self.max_request_size() {
Ok(m) => m,
Err(e) => {
let _ = tx.send(Err(e.into()));
return PrefetchHandle::from(rx);
}
};
if n_elements > max {
let _ = tx.send(Err(CorrelatedStreamError::RequestTooLarge {
requested: n_elements,
max_allowed: max,
}
.into()));
return PrefetchHandle::from(rx);
}
let cmd = Command::Prefetch {
n_elements,
completion: tx,
};
if let Err(e) = self.command_sender.send(cmd) {
if let Command::Prefetch { completion, .. } = e.0 {
let _ = completion.send(Err(CorrelatedStreamError::StreamClosed.into()));
}
}
PrefetchHandle::from(rx)
}
fn position(&self) -> u64 {
self.position.load(Ordering::Acquire)
}
fn buffered(&self) -> u64 {
self.buffered.load(Ordering::Acquire)
}
fn resync(&self, target: u64) -> ResyncHandle<E> {
let (tx, rx) = oneshot::channel();
let cmd = Command::Resync {
target,
completion: tx,
};
if let Err(e) = self.command_sender.send(cmd) {
if let Command::Resync { completion, .. } = e.0 {
let _ = completion.send(Err(CorrelatedStreamError::StreamClosed.into()));
}
}
ResyncHandle::from(rx)
}
}
#[cfg(test)]
mod tests {
use std::{
sync::{
atomic::{AtomicUsize, Ordering},
Arc,
},
time::Duration,
};
use rand::{rngs::StdRng, SeedableRng};
use typenum::U2;
use crate::{
algebra::elliptic_curve::{Curve25519Ristretto, ScalarField},
correlated_randomness::{
generator::CorrelationGenerator,
singlets::{Singlet, Singlets},
stream::{
buffered::{Buffer, BufferConfig, BufferedStream},
errors::CorrelatedStreamError,
CorrelatedStream,
},
},
random::Random,
utils::TryFuture,
};
type Fq = ScalarField<Curve25519Ristretto>;
type TestPB = Singlets<Fq, U2>;
type TestItem = Singlet<Fq>;
type TestErr = CorrelatedStreamError;
#[derive(Clone)]
struct MockGenConfig {
delay: Duration,
fail_after: Option<usize>,
}
impl Default for MockGenConfig {
fn default() -> Self {
Self {
delay: Duration::from_millis(0),
fail_after: None,
}
}
}
struct MockGen {
rng: StdRng,
cfg: MockGenConfig,
items_produced: Arc<AtomicUsize>,
batches: Arc<AtomicUsize>,
}
impl MockGen {
fn new(cfg: MockGenConfig) -> (Self, Arc<AtomicUsize>, Arc<AtomicUsize>) {
let items_produced = Arc::new(AtomicUsize::new(0));
let batches = Arc::new(AtomicUsize::new(0));
let gen = Self {
rng: StdRng::from_seed([0u8; 32]),
cfg,
items_produced: items_produced.clone(),
batches: batches.clone(),
};
(gen, items_produced, batches)
}
}
impl CorrelationGenerator<TestPB> for MockGen {
type Net = ();
type Error = TestErr;
fn run(&mut self, _net: &mut ()) -> impl TryFuture<Ok = TestPB, Error = Self::Error> {
async move { Err(CorrelatedStreamError::StreamClosed) }
}
fn run_for(
&mut self,
n: usize,
_net: &mut (),
) -> impl TryFuture<Ok = Vec<TestItem>, Error = Self::Error> {
async move {
self.batches.fetch_add(1, Ordering::SeqCst);
if self.cfg.delay > Duration::ZERO {
tokio::time::sleep(self.cfg.delay).await;
}
if let Some(threshold) = self.cfg.fail_after {
if self.items_produced.load(Ordering::SeqCst) + n > threshold {
return Err(CorrelatedStreamError::StreamClosed);
}
}
let items: Vec<TestItem> = (0..n)
.map(|_| {
Singlet::<Fq>::random_n::<Vec<_>>(&mut self.rng, 1)
.into_iter()
.next()
.unwrap()
})
.collect();
self.items_produced.fetch_add(n, Ordering::SeqCst);
Ok(items)
}
}
}
fn make_stream(
cfg: MockGenConfig,
buf_cfg: BufferConfig,
) -> (
BufferedStream<TestPB, TestErr>,
Arc<AtomicUsize>,
Arc<AtomicUsize>,
) {
let (gen, produced, batches) = MockGen::new(cfg);
let stream = BufferedStream::<TestPB, TestErr>::new(gen, (), buf_cfg);
(stream, produced, batches)
}
struct SkipMockGen {
rng: StdRng,
produced: Arc<AtomicUsize>,
skipped: Arc<AtomicUsize>,
}
impl SkipMockGen {
fn new() -> (Self, Arc<AtomicUsize>, Arc<AtomicUsize>) {
let produced = Arc::new(AtomicUsize::new(0));
let skipped = Arc::new(AtomicUsize::new(0));
let gen = Self {
rng: StdRng::from_seed([7u8; 32]),
produced: produced.clone(),
skipped: skipped.clone(),
};
(gen, produced, skipped)
}
}
impl CorrelationGenerator<TestPB> for SkipMockGen {
type Net = ();
type Error = TestErr;
const SUPPORTS_UNILATERAL_SKIP: bool = true;
fn run(&mut self, _net: &mut ()) -> impl TryFuture<Ok = TestPB, Error = Self::Error> {
async move { Err(CorrelatedStreamError::StreamClosed) }
}
fn run_for(
&mut self,
n: usize,
_net: &mut (),
) -> impl TryFuture<Ok = Vec<TestItem>, Error = Self::Error> {
async move {
let items: Vec<TestItem> = Singlet::<Fq>::random_n::<Vec<_>>(&mut self.rng, n);
self.produced.fetch_add(n, Ordering::SeqCst);
Ok(items)
}
}
fn skip(
&mut self,
n: usize,
_net: &mut (),
) -> impl TryFuture<Ok = (), Error = Self::Error> {
async move {
self.skipped.fetch_add(n, Ordering::SeqCst);
self.produced.fetch_add(n, Ordering::SeqCst);
Ok(())
}
}
}
struct NoSkipMockGen {
rng: StdRng,
}
impl CorrelationGenerator<TestPB> for NoSkipMockGen {
type Net = ();
type Error = TestErr;
fn run(&mut self, _net: &mut ()) -> impl TryFuture<Ok = TestPB, Error = Self::Error> {
async move { Err(CorrelatedStreamError::StreamClosed) }
}
fn run_for(
&mut self,
n: usize,
_net: &mut (),
) -> impl TryFuture<Ok = Vec<TestItem>, Error = Self::Error> {
async move { Ok(Singlet::<Fq>::random_n::<Vec<_>>(&mut self.rng, n)) }
}
}
#[tokio::test]
async fn next_n_resolves_batch_future() {
let (stream, _, _) = make_stream(MockGenConfig::default(), BufferConfig::eager(16));
let fut = stream.next_n(7).expect("request accepted");
let items = fut.await.expect("batch resolves");
assert_eq!(items.len(), 7);
}
#[tokio::test]
async fn request_too_large_rejected() {
let (stream, _, _) = make_stream(
MockGenConfig::default(),
BufferConfig::eager_with(8, 4), );
match stream.next_n(5) {
Err(CorrelatedStreamError::RequestTooLarge {
requested: 5,
max_allowed: 4,
}) => {}
Ok(_) => panic!("must reject n>max"),
Err(e) => panic!("unexpected error: {e:?}"),
}
}
#[tokio::test]
async fn rate_limit_when_exceeding_capacity() {
let cfg = MockGenConfig {
delay: Duration::from_millis(200),
..Default::default()
};
let (stream, _, _) = make_stream(cfg, BufferConfig::eager(4));
let _f1 = stream.next_n(4).unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
let f2 = stream.next_n(4).unwrap();
let results = futures::future::join_all(f2).await;
assert!(
results
.iter()
.all(|r| matches!(r, Err(CorrelatedStreamError::RateLimitExceeded))),
"expected all four to be rate-limited, got {results:?}"
);
}
#[tokio::test]
async fn prefetch_completes_and_serves_subsequent_requests_quickly() {
let cfg = MockGenConfig {
delay: Duration::from_millis(100),
..Default::default()
};
let (stream, produced, _) = make_stream(cfg, BufferConfig::eager(32));
let handle = stream.prefetch_n(10);
handle.await.expect("prefetch completes");
assert!(produced.load(Ordering::SeqCst) >= 10);
let start = std::time::Instant::now();
let items = stream.next_n(10).unwrap().await.expect("served");
assert_eq!(items.len(), 10);
assert!(
start.elapsed() < Duration::from_millis(80),
"request should be served from prefetched buffer (took {:?})",
start.elapsed()
);
}
#[tokio::test]
async fn sequential_prefetches_all_resolve() {
let (stream, _, _) = make_stream(MockGenConfig::default(), BufferConfig::lazy(16, 0));
for i in 0..2 {
tokio::time::timeout(Duration::from_secs(1), stream.prefetch_n(4))
.await
.unwrap_or_else(|_| panic!("prefetch {i} timed out"))
.expect("prefetch completes");
}
tokio::time::timeout(Duration::from_secs(1), stream.prefetch_n(8))
.await
.expect("partially covered prefetch timed out")
.expect("prefetch completes");
}
#[tokio::test]
async fn generator_error_propagates_to_pending_consumers() {
let cfg = MockGenConfig {
delay: Duration::from_millis(20),
fail_after: Some(0), };
let (stream, _, _) = make_stream(cfg, BufferConfig::eager(16));
let futs = stream.next_n(4).unwrap();
let results = futures::future::join_all(futs).await;
assert!(
results
.iter()
.all(|r| matches!(r, Err(CorrelatedStreamError::StreamClosed))),
"all consumers should receive the generator error"
);
}
#[tokio::test]
async fn fifo_order_across_two_batches() {
let cfg = MockGenConfig {
delay: Duration::from_millis(40),
..Default::default()
};
let (stream, _, batches) = make_stream(cfg, BufferConfig::eager(32));
let f1 = stream.next_n(3).unwrap();
let f2 = stream.next_n(3).unwrap();
let (a, b) = tokio::join!(f1, f2);
let a = a.expect("first batch resolves");
let b = b.expect("second batch resolves");
assert_eq!(a.len(), 3);
assert_eq!(b.len(), 3);
assert!(batches.load(Ordering::SeqCst) >= 1);
}
#[tokio::test]
async fn buffer_config_setters_are_visible() {
let (stream, _, _) = make_stream(MockGenConfig::default(), BufferConfig::eager(16));
assert_eq!(stream.capacity().unwrap(), 16);
stream.set_capacity(32).unwrap();
assert_eq!(stream.capacity().unwrap(), 32);
}
#[tokio::test]
async fn next_n_rejects_too_large() {
let (stream, _, _) = make_stream(
MockGenConfig::default(),
BufferConfig::eager_with(8, 4), );
match stream.next_n(5) {
Err(CorrelatedStreamError::RequestTooLarge {
requested: 5,
max_allowed: 4,
}) => {}
Ok(_) => panic!("must reject n > max"),
Err(e) => panic!("unexpected error: {e:?}"),
}
}
#[tokio::test]
async fn next_n_rate_limit_when_exceeding_capacity() {
let cfg = MockGenConfig {
delay: Duration::from_millis(200),
..Default::default()
};
let (stream, _, _) = make_stream(cfg, BufferConfig::eager(4));
let _f1 = stream.next_n(4).unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
let f2 = stream.next_n(4).unwrap();
assert!(
matches!(f2.await, Err(CorrelatedStreamError::RateLimitExceeded)),
"second next_n should be rate-limited"
);
}
#[tokio::test]
async fn next_n_admitted_when_shortfall_fits_capacity() {
let cfg = MockGenConfig {
delay: Duration::from_millis(50),
..Default::default()
};
let (stream, _, _) = make_stream(cfg, BufferConfig::lazy(8, 0));
let f1 = stream.next_n(4).unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
let f2 = stream.next_n(4).unwrap();
let (a, b) = tokio::join!(f1, f2);
assert_eq!(a.expect("first batch resolves").len(), 4);
assert_eq!(b.expect("second batch resolves").len(), 4);
}
#[tokio::test]
async fn next_n_error_propagates() {
let cfg = MockGenConfig {
delay: Duration::from_millis(20), fail_after: Some(0),
};
let (stream, _, _) = make_stream(cfg, BufferConfig::eager(16));
let result = stream.next_n(4).unwrap().await;
assert!(
matches!(result, Err(CorrelatedStreamError::StreamClosed)),
"expected generator error to propagate through BatchFuture, got {result:?}"
);
}
#[tokio::test]
async fn refill_threshold_drives_proactive_generation() {
let (stream, produced, _) =
make_stream(MockGenConfig::default(), BufferConfig::lazy(16, 8));
let _ = stream.next_n(4).unwrap().await.expect("served");
tokio::time::sleep(Duration::from_millis(50)).await;
let total = produced.load(Ordering::SeqCst);
assert!(total >= 8, "expected >= 8 items generated, got {total}");
}
#[tokio::test]
async fn position_tracks_delivered_not_prefetched() {
let (stream, _, _) = make_stream(MockGenConfig::default(), BufferConfig::eager(32));
stream.prefetch_n(10).await.expect("prefetch completes");
assert_eq!(stream.position(), 0, "prefetch must not advance position");
let _ = stream.next_n(7).unwrap().await.expect("served");
assert_eq!(stream.position(), 7);
let _ = stream.next_n(3).unwrap().await.expect("served");
assert_eq!(stream.position(), 10);
}
#[tokio::test]
async fn buffered_tracks_ready_elements() {
let (stream, _, _) = make_stream(MockGenConfig::default(), BufferConfig::lazy(32, 0));
assert_eq!(stream.buffered(), 0);
stream.prefetch_n(10).await.expect("prefetch completes");
assert_eq!(stream.buffered(), 10, "prefetch fills the buffer");
let _ = stream.next_n(4).unwrap().await.expect("served");
assert_eq!(stream.buffered(), 6, "delivery drains the buffer");
}
#[tokio::test]
async fn position_does_not_advance_on_rejected_request() {
let cfg = MockGenConfig {
delay: Duration::from_millis(200),
..Default::default()
};
let (stream, _, _) = make_stream(cfg, BufferConfig::eager(4));
let _f1 = stream.next_n(4).unwrap();
tokio::time::sleep(Duration::from_millis(20)).await;
let f2 = stream.next_n(4).unwrap();
assert!(matches!(
f2.await,
Err(CorrelatedStreamError::RateLimitExceeded)
));
assert_eq!(stream.position(), 4);
}
#[tokio::test]
async fn resync_drains_buffer_and_advances_position() {
let (stream, produced, _) =
make_stream(MockGenConfig::default(), BufferConfig::lazy(32, 0));
stream.prefetch_n(10).await.expect("prefetch completes");
let before = produced.load(Ordering::SeqCst);
stream.resync(6).await.expect("resync completes");
assert_eq!(stream.position(), 6);
assert_eq!(
produced.load(Ordering::SeqCst),
before,
"drain-only resync must not generate"
);
let _ = stream.next_n(2).unwrap().await.expect("served");
assert_eq!(stream.position(), 8);
}
#[tokio::test]
async fn resync_noop_when_already_at_target() {
let (stream, _, _) = make_stream(MockGenConfig::default(), BufferConfig::eager(16));
let _ = stream.next_n(5).unwrap().await.expect("served");
stream.resync(5).await.expect("no-op resync completes");
assert_eq!(stream.position(), 5);
}
#[tokio::test]
async fn resync_rewind_is_rejected() {
let (stream, _, _) = make_stream(MockGenConfig::default(), BufferConfig::eager(16));
let _ = stream.next_n(5).unwrap().await.expect("served");
match stream.resync(3).await {
Err(CorrelatedStreamError::ResyncRewind {
current: 5,
target: 3,
}) => {}
other => panic!("expected ResyncRewind, got {other:?}"),
}
assert_eq!(stream.position(), 5);
}
#[tokio::test]
async fn resync_unsupported_when_deficit_and_no_skip() {
let gen = NoSkipMockGen {
rng: StdRng::from_seed([0u8; 32]),
};
let stream = BufferedStream::<TestPB, TestErr>::new(gen, (), BufferConfig::lazy(16, 0));
match stream.resync(10).await {
Err(CorrelatedStreamError::ResyncUnsupported {
generated: 0,
target: 10,
}) => {}
other => panic!("expected ResyncUnsupported, got {other:?}"),
}
assert_eq!(stream.position(), 0);
}
#[tokio::test]
async fn resync_skips_at_generator_when_supported() {
let (gen, produced, skipped) = SkipMockGen::new();
let stream = BufferedStream::<TestPB, TestErr>::new(gen, (), BufferConfig::lazy(64, 0));
stream.resync(10).await.expect("resync via skip completes");
assert_eq!(stream.position(), 10);
assert_eq!(
skipped.load(Ordering::SeqCst),
10,
"deficit should be skipped"
);
assert_eq!(
produced.load(Ordering::SeqCst),
10,
"skip advances the generator without delivering items"
);
let items = stream.next_n(3).unwrap().await.expect("served");
assert_eq!(items.len(), 3);
assert_eq!(stream.position(), 13);
}
#[tokio::test]
async fn resync_partial_buffer_then_skip_remainder() {
let (gen, produced, skipped) = SkipMockGen::new();
let stream = BufferedStream::<TestPB, TestErr>::new(gen, (), BufferConfig::lazy(64, 0));
stream.prefetch_n(4).await.expect("prefetch completes");
let produced_after_prefetch = produced.load(Ordering::SeqCst);
assert_eq!(produced_after_prefetch, 4);
stream.resync(10).await.expect("resync completes");
assert_eq!(stream.position(), 10);
assert_eq!(
skipped.load(Ordering::SeqCst),
6,
"only the unbuffered remainder is skipped at the generator"
);
}
}