//! Mock implementations of runtime primitives for testing. #[cfg(any(test, feature = "test-utils"))] pub use crate::storage::memory::Storage as MemoryStorage; use crate::{ Blob, BlobVersion, BufMut, BufferPool, BufferPooler, Clock, Error, Handle, IoBufs, IoBufsMut, Metrics, Name, ReadOptions, Spawner, Storage, Supervisor, WriteOptions, signal::Signal, telemetry::metrics::{Metric, Registered}, }; use bytes::{Bytes, BytesMut}; use commonware_utils::{ channel::{fallible::OneshotExt, oneshot}, sync::Mutex, }; use governor::clock::{Clock as GovernorClock, ReasonablyRealtime}; use rand::{TryCryptoRng, TryRng}; use std::{ future::{Future, poll_fn}, mem, sync::Arc, task::Poll, }; /// Default buffer size (64 KB). Controls both how much data the stream /// pulls per recv and the backpressure threshold for send. const DEFAULT_BUFFER_SIZE: usize = 64 * 1024; /// A mock channel struct that is used internally by Sink and Stream. pub struct Channel { /// Stores the bytes sent by the sink that are not yet read by the stream. buffer: BytesMut, /// If the stream is waiting to read bytes, the waiter stores the number of /// bytes that the stream is waiting for, as well as the oneshot sender that /// the sink uses to send the bytes to the stream directly. waiter: Option<(usize, oneshot::Sender)>, /// Target buffer size, used to bound both the stream's local buffer /// and the shared buffer (backpressure threshold). buffer_size: usize, /// If the sink is blocked waiting for the buffer to drain, this holds /// the oneshot sender that the stream uses to wake the sink. drain_waiter: Option>, /// Tracks whether the sink is still alive and able to send messages. sink_alive: bool, /// Tracks whether the stream is still alive and able to receive messages. stream_alive: bool, } impl Channel { /// Returns an async-safe Sink/Stream pair with default buffer size. pub fn init() -> (Sink, Stream) { Self::init_with_buffer_size(DEFAULT_BUFFER_SIZE) } /// Returns an async-safe Sink/Stream pair with the specified buffer size. pub fn init_with_buffer_size(buffer_size: usize) -> (Sink, Stream) { let channel = Arc::new(Mutex::new(Self { buffer: BytesMut::new(), waiter: None, buffer_size, drain_waiter: None, sink_alive: true, stream_alive: true, })); ( Sink { channel: channel.clone(), state: SinkState::Open, }, Stream { channel, buffer: BytesMut::new(), poisoned: false, }, ) } /// Restores bytes that were detached from the front of the shared buffer. fn restore_front(&mut self, data: Bytes) { if data.is_empty() { return; } let mut restored = BytesMut::with_capacity(data.len() + self.buffer.len()); restored.extend_from_slice(&data); restored.extend_from_slice(&self.buffer); self.buffer = restored; } /// Marks the sink as closed and wakes any waiter. fn close_sink(&mut self) { self.sink_alive = false; // If there is a waiter, resolve it by dropping the oneshot sender. self.waiter.take(); } } struct RecvWaiterGuard { channel: Arc>, active: bool, } impl RecvWaiterGuard { const fn new(channel: Arc>) -> Self { Self { channel, active: true, } } const fn disarm(&mut self) { self.active = false; } } impl Drop for RecvWaiterGuard { fn drop(&mut self) { if !self.active { return; } self.channel.lock().waiter.take(); } } /// A mock sink that implements the Sink trait. pub struct Sink { channel: Arc>, state: SinkState, } /// Lifecycle state for the mock sink half. enum SinkState { /// Sends may be attempted. Open, /// A send is currently in progress. Sending, /// The sink has been closed. Closed, } impl Sink { fn close(&mut self) { if matches!(self.state, SinkState::Closed) { return; } self.channel.lock().close_sink(); self.state = SinkState::Closed; } } impl crate::Sink for Sink { async fn send(&mut self, bufs: impl Into + Send) -> Result<(), Error> { match self.state { SinkState::Open => {} SinkState::Sending => { self.close(); return Err(Error::Closed); } SinkState::Closed => return Err(Error::Closed), } let drain_recv = { let mut channel = self.channel.lock(); // If the receiver is dead, we cannot send any more messages. if !channel.stream_alive { channel.close_sink(); self.state = SinkState::Closed; return Err(Error::SendFailed); } channel.buffer.put(bufs.into()); // If there is a waiter and the buffer is large enough, // resolve the waiter (while clearing the waiter field). if channel .waiter .as_ref() .is_some_and(|(requested, _)| *requested <= channel.buffer.len()) { // Send up to buffer_size bytes (but at least requested amount) let (requested, os_send) = channel.waiter.take().unwrap(); let send_amount = channel.buffer.len().min(requested.max(channel.buffer_size)); let data = channel.buffer.split_to(send_amount).freeze(); // A canceled recv should behave like a buffered transport: // preserve the bytes and allow a subsequent recv to consume them. if let Err(data) = os_send.send(data) { channel.restore_front(data); if !channel.stream_alive { channel.close_sink(); self.state = SinkState::Closed; return Err(Error::SendFailed); } } } // If the buffer exceeds the write limit, block until the // receiver drains enough data. if channel.buffer.len() > channel.buffer_size { assert!(channel.drain_waiter.is_none()); let (os_send, os_recv) = oneshot::channel(); channel.drain_waiter = Some(os_send); os_recv } else { return Ok(()); } }; // Mark the sink as sending before awaiting so cancellation can be // detected by the next send. self.state = SinkState::Sending; // Wait for the receiver to drain the buffer. match drain_recv.await { Ok(()) => { self.state = SinkState::Open; Ok(()) } Err(_) => { self.close(); Err(Error::SendFailed) } } } } impl Drop for Sink { fn drop(&mut self) { self.close(); } } /// A mock stream that implements the Stream trait. pub struct Stream { channel: Arc>, /// Local buffer for data that has been received but not yet consumed. buffer: BytesMut, poisoned: bool, } impl crate::Stream for Stream { async fn recv(&mut self, len: usize) -> Result { if self.poisoned { return Err(Error::Closed); } let os_recv = { let mut channel = self.channel.lock(); // Pull data from channel buffer into local buffer. let target = len.max(channel.buffer_size); let pull_amount = channel .buffer .len() .min(target.saturating_sub(self.buffer.len())); if pull_amount > 0 { let data = channel.buffer.split_to(pull_amount); self.buffer.extend_from_slice(&data); // Wake a blocked sender if the buffer drained below the limit. if channel.buffer.len() <= channel.buffer_size && let Some(sender) = channel.drain_waiter.take() { sender.send_lossy(()); } } // If we have enough, return immediately. if self.buffer.len() >= len { return Ok(IoBufs::from(self.buffer.split_to(len).freeze())); } // If the sink is dead, we cannot receive any more messages. if !channel.sink_alive { self.poisoned = true; return Err(Error::RecvFailed); } // Set up waiter for remaining amount. let remaining = len - self.buffer.len(); assert!(channel.waiter.is_none()); let (os_send, os_recv) = oneshot::channel(); channel.waiter = Some((remaining, os_send)); os_recv }; let mut waiter_guard = RecvWaiterGuard::new(self.channel.clone()); // Pre-poison so that cancellation leaves the stream permanently closed. self.poisoned = true; // Wait for the waiter to be resolved. let data = match os_recv.await { Ok(data) => { waiter_guard.disarm(); self.poisoned = false; data } Err(_) => { waiter_guard.disarm(); return Err(Error::RecvFailed); } }; self.buffer.extend_from_slice(&data); assert!(self.buffer.len() >= len); Ok(IoBufs::from(self.buffer.split_to(len).freeze())) } fn peek(&self, max_len: usize) -> &[u8] { let len = max_len.min(self.buffer.len()); &self.buffer[..len] } } impl Drop for Stream { fn drop(&mut self) { let mut channel = self.channel.lock(); channel.stream_alive = false; // Wake a blocked sender so it can observe the closed stream. channel.drain_waiter.take(); } } /// A sync deferred by a [DelayedSyncBlob], held open until explicitly completed. pub struct DeferredSync { /// Completes the sync with the provided result (success runs the inner blob's sync). pub release: oneshot::Sender>, /// Resolves once the deferred sync's handle begins waiting on `release`. pub blocked: oneshot::Receiver<()>, } /// Coordinates durability operations for a [DelayedSyncContext] or [DelayedSyncBlob]. /// /// Every started sync parks in a deferred queue (in start order) until a test /// releases it or [Self::unblock] runs. [Self::arm] additionally installs a one-shot gate that blocks /// the next durability operation and counts operations from that point on /// ([Self::calls]). The gate is pushed onto the deferred queue when [Self::arm] /// is called, before any operation reaches it. #[derive(Clone, Default)] pub struct PendingSyncs { state: Arc>, } /// State shared by all clones of a [PendingSyncs]. #[derive(Default)] struct State { /// Deferred syncs in start order. syncs: Vec, /// One-shot gate blocking the next durability operation (see [PendingSyncs::arm]). gate: SyncGateState, /// Sticky: stop parking started syncs (see [PendingSyncs::unblock]). unblocked: bool, /// Sticky: syncs resolve to an injected error (see [PendingSyncs::arm_fail]). fail: bool, /// Started syncs issued. starts: usize, /// Started syncs whose completion futures have begun executing. entered: usize, /// Started syncs that completed durably. completions: usize, } impl State { /// Creates a waiter parked in the deferred queue. fn defer(&mut self) -> SyncWaiter { let (release, release_rx) = oneshot::channel(); let (entered, blocked) = oneshot::channel(); self.syncs.push(DeferredSync { release, blocked }); SyncWaiter { entered, release: release_rx, } } /// Records a durability operation if the gate is armed, returning the /// one-shot gate waiter if it has not been consumed yet. const fn observe(&mut self) -> Option { if !self.gate.tracking { return None; } self.gate.calls += 1; self.gate.waiter.take() } /// Parks a new deferred sync, unless [PendingSyncs::unblock] already ran. fn park(&mut self) -> Option { if self.unblocked { return None; } Some(self.defer()) } } /// Forwards [Supervisor], [Clock], [GovernorClock], [ReasonablyRealtime], /// [Metrics], [BufferPooler], [TryRng], and [TryCryptoRng] to the wrapped /// context for test context wrappers with one extra field (named by the /// second argument). macro_rules! forward_context { ($wrapper:ident, $field:ident) => { impl Supervisor for $wrapper { fn name(&self) -> Name { self.inner.name() } fn child(&self, label: &'static str) -> Self { Self { inner: self.inner.child(label), $field: self.$field.clone(), } } fn with_attribute(self, key: &'static str, value: impl std::fmt::Display) -> Self { Self { inner: self.inner.with_attribute(key, value), $field: self.$field, } } } impl Clock for $wrapper { fn current(&self) -> std::time::SystemTime { self.inner.current() } fn sleep( &self, duration: std::time::Duration, ) -> impl Future + Send + 'static { self.inner.sleep(duration) } fn sleep_until( &self, deadline: std::time::SystemTime, ) -> impl Future + Send + 'static { self.inner.sleep_until(deadline) } } impl GovernorClock for $wrapper { type Instant = std::time::SystemTime; fn now(&self) -> Self::Instant { self.current() } } impl ReasonablyRealtime for $wrapper {} impl Metrics for $wrapper { fn register, H: Into, M: Metric>( &self, name: N, help: H, metric: M, ) -> Registered { self.inner.register(name, help, metric) } fn encode(&self) -> String { self.inner.encode() } } impl BufferPooler for $wrapper { fn network_buffer_pool(&self) -> &BufferPool { self.inner.network_buffer_pool() } fn storage_buffer_pool(&self) -> &BufferPool { self.inner.storage_buffer_pool() } } impl TryRng for $wrapper { type Error = E::Error; fn try_next_u32(&mut self) -> Result { self.inner.try_next_u32() } fn try_next_u64(&mut self) -> Result { self.inner.try_next_u64() } fn try_fill_bytes(&mut self, dest: &mut [u8]) -> Result<(), Self::Error> { self.inner.try_fill_bytes(dest) } } impl TryCryptoRng for $wrapper {} }; } /// Snapshot of the options observed by a [RecordingContext] or [RecordingBlob]. #[cfg(any(test, feature = "test-utils"))] #[derive(Clone, Debug, Default, Eq, PartialEq)] pub struct RecordingSnapshot { /// Options supplied to read operations, in call order. pub reads: Vec, /// Options supplied to write operations, in call order. pub writes: Vec, } /// Shared observations produced by recording storage wrappers. #[cfg(any(test, feature = "test-utils"))] #[derive(Clone, Default)] pub struct Recordings { state: Arc>, } #[cfg(any(test, feature = "test-utils"))] impl Recordings { /// Return a snapshot of all observations recorded so far. pub fn snapshot(&self) -> RecordingSnapshot { self.state.lock().clone() } /// Remove all recorded observations. pub fn clear(&self) { *self.state.lock() = RecordingSnapshot::default(); } fn read(&self, options: ReadOptions) { self.state.lock().reads.push(options); } fn write(&self, options: WriteOptions) { self.state.lock().writes.push(options); } } /// Context wrapper that records options supplied to every opened blob. #[cfg(any(test, feature = "test-utils"))] #[derive(Clone)] pub struct RecordingContext { /// Wrapped context. pub inner: E, /// Observations shared by this context and all blobs opened through it. pub recordings: Recordings, } #[cfg(any(test, feature = "test-utils"))] impl RecordingContext { /// Wrap `inner` and return both the context and its shared observations. pub fn new(inner: E) -> (Self, Recordings) { let recordings = Recordings::default(); ( Self { inner, recordings: recordings.clone(), }, recordings, ) } } #[cfg(any(test, feature = "test-utils"))] forward_context!(RecordingContext, recordings); #[cfg(any(test, feature = "test-utils"))] impl Spawner for RecordingContext { fn shared(mut self, blocking: bool) -> Self { self.inner = self.inner.shared(blocking); self } fn dedicated(mut self) -> Self { self.inner = self.inner.dedicated(); self } fn spawn(self, f: F) -> Handle where F: FnOnce(Self) -> Fut + Send + 'static, Fut: Future + Send + 'static, T: Send + 'static, { let recordings = self.recordings; self.inner.spawn(move |inner| f(Self { inner, recordings })) } async fn stop(self, value: i32, timeout: Option) -> Result<(), Error> { self.inner.stop(value, timeout).await } fn stopped(&self) -> Signal { self.inner.stopped() } } #[cfg(any(test, feature = "test-utils"))] impl Storage for RecordingContext { type Blob = RecordingBlob; async fn open_versioned( &self, partition: &str, name: &[u8], versions: std::ops::RangeInclusive, ) -> Result<(Self::Blob, u64, BlobVersion), Error> { let (inner, len, version) = self.inner.open_versioned(partition, name, versions).await?; Ok(( RecordingBlob { inner, recordings: self.recordings.clone(), }, len, version, )) } async fn remove(&self, partition: &str, name: Option<&[u8]>) -> Result<(), Error> { self.inner.remove(partition, name).await } async fn scan(&self, partition: &str) -> Result>, Error> { self.inner.scan(partition).await } } /// Blob wrapper that records read and write options before delegating each operation. #[cfg(any(test, feature = "test-utils"))] #[derive(Clone)] pub struct RecordingBlob { inner: B, recordings: Recordings, } #[cfg(any(test, feature = "test-utils"))] impl Blob for RecordingBlob { async fn read_at_buf( &self, offset: u64, len: usize, bufs: impl Into + Send, options: ReadOptions, ) -> Result { self.recordings.read(options); self.inner.read_at_buf(offset, len, bufs, options).await } async fn read_at( &self, offset: u64, len: usize, options: ReadOptions, ) -> Result { self.recordings.read(options); self.inner.read_at(offset, len, options).await } async fn write_at( &self, offset: u64, bufs: impl Into + Send, options: WriteOptions, ) -> Result<(), Error> { self.recordings.write(options); self.inner.write_at(offset, bufs, options).await } async fn resize(&self, len: u64) -> Result<(), Error> { self.inner.resize(len).await } async fn sync(&self) -> Result<(), Error> { self.inner.sync().await } async fn start_sync(&self) -> Handle<()> { self.inner.start_sync().await } } /// Context wrapper whose blobs defer [Blob::start_sync] and can gate blocking syncs in tests. #[derive(Clone)] pub struct DelayedSyncContext { pub inner: E, pub pending: PendingSyncs, } forward_context!(DelayedSyncContext, pending); impl Spawner for DelayedSyncContext { fn shared(mut self, blocking: bool) -> Self { self.inner = self.inner.shared(blocking); self } fn dedicated(mut self) -> Self { self.inner = self.inner.dedicated(); self } fn spawn(self, f: F) -> Handle where F: FnOnce(Self) -> Fut + Send + 'static, Fut: Future + Send + 'static, T: Send + 'static, { let pending = self.pending; self.inner.spawn(move |inner| f(Self { inner, pending })) } async fn stop(self, value: i32, timeout: Option) -> Result<(), Error> { self.inner.stop(value, timeout).await } fn stopped(&self) -> Signal { self.inner.stopped() } } impl Storage for DelayedSyncContext { type Blob = DelayedSyncBlob; async fn open_versioned( &self, partition: &str, name: &[u8], versions: std::ops::RangeInclusive, ) -> Result<(Self::Blob, u64, BlobVersion), Error> { let (inner, len, version) = self.inner.open_versioned(partition, name, versions).await?; Ok(( DelayedSyncBlob { inner, pending: self.pending.clone(), }, len, version, )) } async fn remove(&self, partition: &str, name: Option<&[u8]>) -> Result<(), Error> { self.inner.remove(partition, name).await } async fn scan(&self, partition: &str) -> Result>, Error> { self.inner.scan(partition).await } } /// Blob wrapper that parks each started sync and supports one-shot blocking sync tracking. #[derive(Clone)] pub struct DelayedSyncBlob { inner: B, pending: PendingSyncs, } impl DelayedSyncBlob { /// Wrap `inner`, returning the blob and the list its deferred syncs are pushed onto. pub fn new(inner: B) -> (Self, PendingSyncs) { let pending = PendingSyncs::default(); ( Self { inner, pending: pending.clone(), }, pending, ) } } impl Blob for DelayedSyncBlob { async fn read_at_buf( &self, offset: u64, len: usize, bufs: impl Into + Send, options: ReadOptions, ) -> Result { self.inner.read_at_buf(offset, len, bufs, options).await } async fn read_at( &self, offset: u64, len: usize, options: ReadOptions, ) -> Result { self.inner.read_at(offset, len, options).await } async fn write_at( &self, offset: u64, bufs: impl Into + Send, options: WriteOptions, ) -> Result<(), Error> { if !options.contains(WriteOptions::SYNC) || !self.pending.tracking() { return self.inner.write_at(offset, bufs, options).await; } self.inner .write_at(offset, bufs, options.without(WriteOptions::SYNC)) .await?; self.sync().await } async fn resize(&self, len: u64) -> Result<(), Error> { self.inner.resize(len).await } async fn sync(&self) -> Result<(), Error> { self.pending.wait().await?; self.inner.sync().await } async fn start_sync(&self) -> Handle<()> { let pending = self.pending.clone(); let inner = self.inner.clone(); let waiter = { let mut state = pending.state.lock(); state.starts += 1; // An armed gate takes precedence over parking. state.observe().or_else(|| state.park()) }; Handle::from_future(async move { let fail = { let mut state = pending.state.lock(); state.entered += 1; state.fail }; match waiter { Some(waiter) => waiter.wait().await?, None if fail => return Err(injected_sync_failure()), None => {} } inner.sync().await?; pending.state.lock().completions += 1; Ok(()) }) } } /// Take the oldest pending sync, panicking if none was started. pub fn next_pending_sync(pending: &PendingSyncs) -> DeferredSync { let mut pending = pending.lock(); assert!(!pending.is_empty(), "no pending sync was started"); pending.remove(0) } /// Complete the oldest `count` pending syncs successfully. pub fn release_next_pending_syncs(pending: &PendingSyncs, count: usize) { let syncs = { let mut pending = pending.lock(); assert!( pending.len() >= count, "not enough pending syncs: have {}, need {count}", pending.len() ); pending.drain(..count).collect::>() }; for sync in syncs { let _ = sync.release.send(Ok(())); } } /// Complete all pending syncs successfully. pub fn release_pending_syncs(pending: &PendingSyncs) { for sync in mem::take(&mut *pending.lock()) { let _ = sync.release.send(Ok(())); } } /// Drive `fut` to completion, releasing any parked syncs each time it stalls. pub async fn drive_pending_syncs(pending: &PendingSyncs, fut: impl Future) -> T { let mut fut = std::pin::pin!(fut); poll_fn(|cx| match fut.as_mut().poll(cx) { Poll::Ready(out) => Poll::Ready(out), Poll::Pending => { // A concurrent task may park a new sync after this release, so // self-wake to check again on the next scheduler tick. release_pending_syncs(pending); cx.waker().wake_by_ref(); Poll::Pending } }) .await } /// Fail all pending syncs with an injected I/O error. pub fn fail_pending_syncs(pending: &PendingSyncs) { for sync in mem::take(&mut *pending.lock()) { let _ = sync.release.send(Err(injected_sync_failure())); } } /// The error injected by [fail_pending_syncs] and [PendingSyncs::arm_fail]. fn injected_sync_failure() -> Error { Error::Io(std::io::Error::other("injected sync failure").into()) } struct SyncWaiter { entered: oneshot::Sender<()>, release: oneshot::Receiver>, } impl SyncWaiter { async fn wait(self) -> Result<(), Error> { self.entered.send_lossy(()); self.release.await.map_err(|_| Error::Closed)??; Ok(()) } } #[derive(Default)] struct SyncGateState { tracking: bool, calls: usize, waiter: Option, } impl PendingSyncs { /// Locks the deferred sync queue. pub fn lock(&self) -> commonware_utils::sync::MappedMutexGuard<'_, Vec> { commonware_utils::sync::MutexGuard::map(self.state.lock(), |state| &mut state.syncs) } /// Begins counting durability operations and blocks the next one behind a /// one-shot gate (pushed onto the deferred queue so tests can release it). /// /// Once the gate is consumed, started syncs park in the deferred queue as /// usual while [Self::calls] keeps counting. pub fn arm(&self) { let mut state = self.state.lock(); assert!(!state.gate.tracking, "sync gate already armed"); assert!( state.gate.waiter.is_none(), "sync gate already has a waiter" ); state.gate.tracking = true; state.gate.calls = 0; let waiter = state.defer(); state.gate.waiter = Some(waiter); } /// Returns the number of durability operations observed since [Self::arm]. pub fn calls(&self) -> usize { self.state.lock().gate.calls } fn tracking(&self) -> bool { self.state.lock().gate.tracking } /// Releases every parked sync and permanently stops parking new ones: future started syncs /// proceed immediately (or fail, after [Self::arm_fail]). Unlike [release_pending_syncs], /// which drains the queue once, this is sticky. A gate armed via [Self::arm] parks in the /// same queue, so this releases it too. pub fn unblock(&self) { let (drained, fail) = { let mut state = self.state.lock(); state.unblocked = true; (mem::take(&mut state.syncs), state.fail) }; for sync in drained { let result = if fail { Err(injected_sync_failure()) } else { Ok(()) }; let _ = sync.release.send(result); } } /// Arms every parked and future started sync to resolve to an injected error once released /// via [Self::unblock]. The release helpers send explicit results and ignore this. pub fn arm_fail(&self) { self.state.lock().fail = true; } /// Number of started syncs issued. pub fn starts(&self) -> usize { self.state.lock().starts } /// Number of started syncs whose completion futures have begun executing, parked or not. pub fn entered(&self) -> usize { self.state.lock().entered } /// Number of started syncs that completed durably. pub fn completions(&self) -> usize { self.state.lock().completions } async fn wait(&self) -> Result<(), Error> { let waiter = self.state.lock().observe(); match waiter { Some(waiter) => waiter.wait().await, None => Ok(()), } } } /// Controls a [WriteFaultContext]: while armed, every `write_at` fails with an /// injected error. Successful writes are counted. #[derive(Clone, Default)] pub struct WriteFaults { state: Arc>, } #[derive(Default)] struct WriteFaultState { fail: bool, writes: u64, } impl WriteFaults { /// Start failing writes. pub fn arm(&self) { self.state.lock().fail = true; } /// Stop failing writes. pub fn disarm(&self) { self.state.lock().fail = false; } /// The number of successful writes so far. pub fn writes(&self) -> u64 { self.state.lock().writes } fn check(&self) -> Result<(), Error> { if self.state.lock().fail { return Err(Error::Io( std::io::Error::other("injected write failure").into(), )); } Ok(()) } fn note(&self) { self.state.lock().writes += 1; } } /// Context wrapper whose blobs fail `write_at` while the shared [WriteFaults] /// is armed, counting successful writes. Unlike [DelayedSyncContext], this injects failures /// into inline writes issued before any blob sync starts. #[derive(Clone)] pub struct WriteFaultContext { pub inner: E, pub faults: WriteFaults, } forward_context!(WriteFaultContext, faults); impl Storage for WriteFaultContext { type Blob = WriteFaultBlob; async fn open_versioned( &self, partition: &str, name: &[u8], versions: std::ops::RangeInclusive, ) -> Result<(Self::Blob, u64, BlobVersion), Error> { let (inner, len, version) = self.inner.open_versioned(partition, name, versions).await?; Ok(( WriteFaultBlob { inner, faults: self.faults.clone(), }, len, version, )) } async fn remove(&self, partition: &str, name: Option<&[u8]>) -> Result<(), Error> { self.inner.remove(partition, name).await } async fn scan(&self, partition: &str) -> Result>, Error> { self.inner.scan(partition).await } } /// Blob wrapper that fails `write_at` while its [WriteFaults] is armed. #[derive(Clone)] pub struct WriteFaultBlob { inner: B, faults: WriteFaults, } impl Blob for WriteFaultBlob { async fn read_at_buf( &self, offset: u64, len: usize, bufs: impl Into + Send, options: ReadOptions, ) -> Result { self.inner.read_at_buf(offset, len, bufs, options).await } async fn read_at( &self, offset: u64, len: usize, options: ReadOptions, ) -> Result { self.inner.read_at(offset, len, options).await } async fn write_at( &self, offset: u64, bufs: impl Into + Send, options: WriteOptions, ) -> Result<(), Error> { self.faults.check()?; self.inner.write_at(offset, bufs, options).await?; self.faults.note(); Ok(()) } async fn resize(&self, len: u64) -> Result<(), Error> { self.inner.resize(len).await } async fn sync(&self) -> Result<(), Error> { self.inner.sync().await } async fn start_sync(&self) -> Handle<()> { self.inner.start_sync().await } } /// Context wrapper whose blobs fail `sync` and `start_sync` for a single partition. #[derive(Clone)] pub struct SyncFaultContext { pub inner: E, pub fail_partition: String, } forward_context!(SyncFaultContext, fail_partition); impl Storage for SyncFaultContext { type Blob = SyncFaultBlob; async fn open_versioned( &self, partition: &str, name: &[u8], versions: std::ops::RangeInclusive, ) -> Result<(Self::Blob, u64, BlobVersion), Error> { let (inner, len, version) = self.inner.open_versioned(partition, name, versions).await?; Ok(( SyncFaultBlob { inner, faulty: partition == self.fail_partition, }, len, version, )) } async fn remove(&self, partition: &str, name: Option<&[u8]>) -> Result<(), Error> { self.inner.remove(partition, name).await } async fn scan(&self, partition: &str) -> Result>, Error> { self.inner.scan(partition).await } } /// Blob wrapper that fails `sync` and `start_sync` when marked faulty. #[derive(Clone)] pub struct SyncFaultBlob { inner: B, faulty: bool, } impl Blob for SyncFaultBlob { async fn read_at_buf( &self, offset: u64, len: usize, bufs: impl Into + Send, options: ReadOptions, ) -> Result { self.inner.read_at_buf(offset, len, bufs, options).await } async fn read_at( &self, offset: u64, len: usize, options: ReadOptions, ) -> Result { self.inner.read_at(offset, len, options).await } async fn write_at( &self, offset: u64, bufs: impl Into + Send, options: WriteOptions, ) -> Result<(), Error> { self.inner.write_at(offset, bufs, options).await } async fn resize(&self, len: u64) -> Result<(), Error> { self.inner.resize(len).await } async fn sync(&self) -> Result<(), Error> { if self.faulty { let err = std::io::Error::other("injected partition sync fault"); return Err(Error::Io(err.into())); } self.inner.sync().await } async fn start_sync(&self) -> Handle<()> { if self.faulty { return Handle::ready(self.sync().await); } self.inner.start_sync().await } } #[cfg(test)] mod tests { use super::*; use crate::{Clock, IoBufMut, Runner, Sink, Spawner, Stream, deterministic}; use commonware_macros::select; use std::{thread::sleep, time::Duration}; #[test] fn recording_context_preserves_data_and_records_options() { deterministic::Runner::default().start(|context| async move { let (context, recordings) = RecordingContext::new(context); let (blob, _) = context.open("recording", b"blob").await.unwrap(); blob.write_at(0, b"data", WriteOptions::DONT_CACHE) .await .unwrap(); let read = blob.read_at(0, 4, ReadOptions::DONT_CACHE).await.unwrap(); assert_eq!(read.coalesce(), b"data"); let read = blob .read_at_buf(0, 4, IoBufMut::with_capacity(4), ReadOptions::default()) .await .unwrap(); assert_eq!(read.coalesce(), b"data"); assert_eq!( recordings.snapshot(), RecordingSnapshot { reads: vec![ReadOptions::DONT_CACHE, ReadOptions::default()], writes: vec![WriteOptions::DONT_CACHE], } ); recordings.clear(); assert_eq!(recordings.snapshot(), RecordingSnapshot::default()); }); } async fn assert_read_options_forwarded( context: &E, recordings: &Recordings, partition: &str, ) { let (blob, _) = context.open(partition, b"blob").await.unwrap(); blob.write_at(0, b"data", WriteOptions::default()) .await .unwrap(); recordings.clear(); let read = blob.read_at(0, 4, ReadOptions::DONT_CACHE).await.unwrap(); assert_eq!(read.coalesce(), b"data"); let read = blob .read_at_buf(0, 4, IoBufMut::with_capacity(4), ReadOptions::DONT_CACHE) .await .unwrap(); assert_eq!(read.coalesce(), b"data"); assert_eq!( recordings.snapshot(), RecordingSnapshot { reads: vec![ReadOptions::DONT_CACHE, ReadOptions::DONT_CACHE], writes: Vec::new(), } ); } #[test] fn delayed_sync_blob_forwards_read_options() { deterministic::Runner::default().start(|context| async move { let (inner, recordings) = RecordingContext::new(context); let context = DelayedSyncContext { inner, pending: PendingSyncs::default(), }; assert_read_options_forwarded(&context, &recordings, "delayed_sync").await; }); } #[test] fn write_fault_blob_forwards_read_options() { deterministic::Runner::default().start(|context| async move { let (inner, recordings) = RecordingContext::new(context); let context = WriteFaultContext { inner, faults: WriteFaults::default(), }; assert_read_options_forwarded(&context, &recordings, "write_fault").await; }); } #[test] fn sync_fault_blob_forwards_read_options() { deterministic::Runner::default().start(|context| async move { let (inner, recordings) = RecordingContext::new(context); let context = SyncFaultContext { inner, fail_partition: "sync_fault".to_string(), }; assert_read_options_forwarded(&context, &recordings, "sync_fault").await; }); } #[test] fn test_send_recv() { let (mut sink, mut stream) = Channel::init(); let data = b"hello world"; let executor = deterministic::Runner::default(); executor.start(|_| async move { sink.send(data.as_slice()).await.unwrap(); let received = stream.recv(data.len()).await.unwrap(); assert_eq!(received.coalesce(), data); }); } #[test] fn test_send_recv_partial_multiple() { let (mut sink, mut stream) = Channel::init(); let data = b"hello"; let data2 = b" world"; let executor = deterministic::Runner::default(); executor.start(|_| async move { sink.send(data.as_slice()).await.unwrap(); sink.send(data2.as_slice()).await.unwrap(); let received = stream.recv(5).await.unwrap(); assert_eq!(received.coalesce(), b"hello"); let received = stream.recv(5).await.unwrap(); assert_eq!(received.coalesce(), b" worl"); let received = stream.recv(1).await.unwrap(); assert_eq!(received.coalesce(), b"d"); }); } #[test] fn test_send_recv_async() { let (mut sink, mut stream) = Channel::init(); let data = b"hello world"; let executor = deterministic::Runner::default(); executor.start(|_| async move { let (received, _) = futures::try_join!(stream.recv(data.len()), async { sleep(Duration::from_millis(50)); sink.send(data.as_slice()).await }) .unwrap(); assert_eq!(received.coalesce(), data); }); } #[test] fn test_recv_error_sink_dropped_while_waiting() { let (sink, mut stream) = Channel::init(); let executor = deterministic::Runner::default(); executor.start(|context| async move { futures::join!( async { let result = stream.recv(5).await; assert!(matches!(result, Err(Error::RecvFailed))); let result = stream.recv(5).await; assert!(matches!(result, Err(Error::Closed))); }, async { // Wait for the stream to start waiting context.sleep(Duration::from_millis(50)).await; drop(sink); } ); }); } #[test] fn test_recv_error_sink_dropped_before_recv() { let (sink, mut stream) = Channel::init(); drop(sink); // Drop sink immediately let executor = deterministic::Runner::default(); executor.start(|_| async move { let result = stream.recv(5).await; assert!(matches!(result, Err(Error::RecvFailed))); let result = stream.recv(5).await; assert!(matches!(result, Err(Error::Closed))); }); } #[test] fn test_send_error_stream_dropped() { let (mut sink, mut stream) = Channel::init(); let executor = deterministic::Runner::default(); executor.start(|context| async move { // Send some bytes assert!(sink.send(b"7 bytes".as_slice()).await.is_ok()); // Spawn a task to initiate recv's where the first one will succeed and then will drop. let handle = context.child("recv").spawn(|_| async move { let _ = stream.recv(5).await; let _ = stream.recv(5).await; }); // Give the async task a moment to start context.sleep(Duration::from_millis(50)).await; // Drop the stream by aborting the handle handle.abort(); assert!(matches!(handle.await, Err(Error::Closed))); // Try to send a message. The stream is dropped, so this should fail. let result = sink.send(b"hello world".as_slice()).await; assert!(matches!(result, Err(Error::SendFailed))); let result = sink.send(b"hello world".as_slice()).await; assert!(matches!(result, Err(Error::Closed))); }); } #[test] fn test_send_error_stream_dropped_before_send() { let (mut sink, stream) = Channel::init(); drop(stream); // Drop stream immediately let executor = deterministic::Runner::default(); executor.start(|_| async move { let result = sink.send(b"hello world".as_slice()).await; assert!(matches!(result, Err(Error::SendFailed))); let result = sink.send(b"hello world".as_slice()).await; assert!(matches!(result, Err(Error::Closed))); }); } #[test] fn test_recv_timeout() { let (_sink, mut stream) = Channel::init(); // If there is no data to read, test that the recv function just blocks. // The timeout should return first. let executor = deterministic::Runner::default(); executor.start(|context| async move { select! { v = stream.recv(5) => { panic!("unexpected value: {v:?}"); }, _ = context.sleep(Duration::from_millis(100)) => "timeout", }; }); } #[test] fn test_peek_empty() { let (_sink, stream) = Channel::init(); // Peek on a fresh stream should return empty slice assert!(stream.peek(10).is_empty()); } #[test] fn test_peek_after_partial_recv() { let (mut sink, mut stream) = Channel::init(); let executor = deterministic::Runner::default(); executor.start(|_| async move { // Send more data than we'll consume sink.send(b"hello world".as_slice()).await.unwrap(); // Recv only part of it let received = stream.recv(5).await.unwrap(); assert_eq!(received.coalesce(), b"hello"); // Peek should show the remaining data assert_eq!(stream.peek(100), b" world"); // Peek with smaller max_len assert_eq!(stream.peek(3), b" wo"); // Peek doesn't consume - can peek again assert_eq!(stream.peek(100), b" world"); // Recv consumes the peeked data let received = stream.recv(6).await.unwrap(); assert_eq!(received.coalesce(), b" world"); // Peek is now empty assert!(stream.peek(100).is_empty()); }); } #[test] fn test_peek_after_recv_wakeup() { let (mut sink, mut stream) = Channel::init_with_buffer_size(64); let executor = deterministic::Runner::default(); executor.start(|context| async move { // Spawn recv that will block waiting let (tx, rx) = oneshot::channel(); let recv_handle = context.child("recv").spawn(|_| async move { let data = stream.recv(3).await.unwrap(); tx.send(stream).ok(); data }); // Let recv set up waiter context.sleep(Duration::from_millis(10)).await; // Send more than requested sink.send(b"ABCDEFGHIJ".as_slice()).await.unwrap(); // Recv gets its 3 bytes let received = recv_handle.await.unwrap(); assert_eq!(received.coalesce(), b"ABC"); // Get stream back and verify peek sees remaining data let stream = rx.await.unwrap(); assert_eq!(stream.peek(100), b"DEFGHIJ"); }); } #[test] fn test_peek_multiple_sends() { let (mut sink, mut stream) = Channel::init(); let executor = deterministic::Runner::default(); executor.start(|_| async move { // Send multiple chunks sink.send(b"aaa".as_slice()).await.unwrap(); sink.send(b"bbb".as_slice()).await.unwrap(); sink.send(b"ccc".as_slice()).await.unwrap(); // Recv less than total let received = stream.recv(4).await.unwrap(); assert_eq!(received.coalesce(), b"aaab"); // Peek should show remaining assert_eq!(stream.peek(100), b"bbccc"); }); } #[test] fn test_buffer_size_limit() { // Use a small buffer capacity for testing let (mut sink, mut stream) = Channel::init_with_buffer_size(10); let executor = deterministic::Runner::default(); executor.start(|context| async move { // Send more than buffer capacity concurrently with recv // so the sender can drain via backpressure. let send_handle = context.child("sender").spawn(|_| async move { sink.send(b"0123456789ABCDEF".as_slice()).await.unwrap(); sink }); // Recv a small amount - should only pull up to capacity (10 bytes) let received = stream.recv(2).await.unwrap(); assert_eq!(received.coalesce(), b"01"); // Peek should show remaining buffered data (8 bytes, not 14) assert_eq!(stream.peek(100), b"23456789"); // The rest should still be in the channel buffer // Recv more to pull the remaining data let received = stream.recv(8).await.unwrap(); assert_eq!(received.coalesce(), b"23456789"); // Now peek should show next chunk from channel (up to capacity) let received = stream.recv(2).await.unwrap(); assert_eq!(received.coalesce(), b"AB"); assert_eq!(stream.peek(100), b"CDEF"); // Ensure the sender completes send_handle.await.unwrap(); }); } #[test] fn test_recv_before_send() { // Use a small buffer capacity for testing let (mut sink, mut stream) = Channel::init_with_buffer_size(10); let executor = deterministic::Runner::default(); executor.start(|context| async move { // Start recv before send (will wait) let recv_handle = context .child("recv") .spawn(|_| async move { stream.recv(3).await.unwrap() }); // Give recv time to set up waiter context.sleep(Duration::from_millis(10)).await; // Send more than capacity sink.send(b"ABCDEFGHIJKLMNOP".as_slice()).await.unwrap(); // Recv should get its 3 bytes let received = recv_handle.await.unwrap(); assert_eq!(received.coalesce(), b"ABC"); }); } }