#[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,
};
const DEFAULT_BUFFER_SIZE: usize = 64 * 1024;
pub struct Channel {
buffer: BytesMut,
waiter: Option<(usize, oneshot::Sender<Bytes>)>,
buffer_size: usize,
drain_waiter: Option<oneshot::Sender<()>>,
sink_alive: bool,
stream_alive: bool,
}
impl Channel {
pub fn init() -> (Sink, Stream) {
Self::init_with_buffer_size(DEFAULT_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,
},
)
}
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;
}
fn close_sink(&mut self) {
self.sink_alive = false;
self.waiter.take();
}
}
struct RecvWaiterGuard {
channel: Arc<Mutex<Channel>>,
active: bool,
}
impl RecvWaiterGuard {
const fn new(channel: Arc<Mutex<Channel>>) -> 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();
}
}
pub struct Sink {
channel: Arc<Mutex<Channel>>,
state: SinkState,
}
enum SinkState {
Open,
Sending,
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<IoBufs> + 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 !channel.stream_alive {
channel.close_sink();
self.state = SinkState::Closed;
return Err(Error::SendFailed);
}
channel.buffer.put(bufs.into());
if channel
.waiter
.as_ref()
.is_some_and(|(requested, _)| *requested <= channel.buffer.len())
{
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();
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 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(());
}
};
self.state = SinkState::Sending;
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();
}
}
pub struct Stream {
channel: Arc<Mutex<Channel>>,
buffer: BytesMut,
poisoned: bool,
}
impl crate::Stream for Stream {
async fn recv(&mut self, len: usize) -> Result<IoBufs, Error> {
if self.poisoned {
return Err(Error::Closed);
}
let os_recv = {
let mut channel = self.channel.lock();
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);
if channel.buffer.len() <= channel.buffer_size
&& let Some(sender) = channel.drain_waiter.take()
{
sender.send_lossy(());
}
}
if self.buffer.len() >= len {
return Ok(IoBufs::from(self.buffer.split_to(len).freeze()));
}
if !channel.sink_alive {
self.poisoned = true;
return Err(Error::RecvFailed);
}
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());
self.poisoned = true;
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;
channel.drain_waiter.take();
}
}
pub struct DeferredSync {
pub release: oneshot::Sender<Result<(), Error>>,
pub blocked: oneshot::Receiver<()>,
}
#[derive(Clone, Default)]
pub struct PendingSyncs {
state: Arc<Mutex<State>>,
}
#[derive(Default)]
struct State {
syncs: Vec<DeferredSync>,
gate: SyncGateState,
unblocked: bool,
fail: bool,
starts: usize,
entered: usize,
completions: usize,
}
impl State {
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,
}
}
const fn observe(&mut self) -> Option<SyncWaiter> {
if !self.gate.tracking {
return None;
}
self.gate.calls += 1;
self.gate.waiter.take()
}
fn park(&mut self) -> Option<SyncWaiter> {
if self.unblocked {
return None;
}
Some(self.defer())
}
}
macro_rules! forward_context {
($wrapper:ident, $field:ident) => {
impl<E: Supervisor> Supervisor for $wrapper<E> {
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<E: Clock> Clock for $wrapper<E> {
fn current(&self) -> std::time::SystemTime {
self.inner.current()
}
fn sleep(
&self,
duration: std::time::Duration,
) -> impl Future<Output = ()> + Send + 'static {
self.inner.sleep(duration)
}
fn sleep_until(
&self,
deadline: std::time::SystemTime,
) -> impl Future<Output = ()> + Send + 'static {
self.inner.sleep_until(deadline)
}
}
impl<E: Clock> GovernorClock for $wrapper<E> {
type Instant = std::time::SystemTime;
fn now(&self) -> Self::Instant {
self.current()
}
}
impl<E: Clock> ReasonablyRealtime for $wrapper<E> {}
impl<E: Metrics> Metrics for $wrapper<E> {
fn register<N: Into<String>, H: Into<String>, M: Metric>(
&self,
name: N,
help: H,
metric: M,
) -> Registered<M> {
self.inner.register(name, help, metric)
}
fn encode(&self) -> String {
self.inner.encode()
}
}
impl<E: BufferPooler> BufferPooler for $wrapper<E> {
fn network_buffer_pool(&self) -> &BufferPool {
self.inner.network_buffer_pool()
}
fn storage_buffer_pool(&self) -> &BufferPool {
self.inner.storage_buffer_pool()
}
}
impl<E: TryRng> TryRng for $wrapper<E> {
type Error = E::Error;
fn try_next_u32(&mut self) -> Result<u32, Self::Error> {
self.inner.try_next_u32()
}
fn try_next_u64(&mut self) -> Result<u64, Self::Error> {
self.inner.try_next_u64()
}
fn try_fill_bytes(&mut self, dest: &mut [u8]) -> Result<(), Self::Error> {
self.inner.try_fill_bytes(dest)
}
}
impl<E: TryCryptoRng> TryCryptoRng for $wrapper<E> {}
};
}
#[cfg(any(test, feature = "test-utils"))]
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct RecordingSnapshot {
pub reads: Vec<ReadOptions>,
pub writes: Vec<WriteOptions>,
}
#[cfg(any(test, feature = "test-utils"))]
#[derive(Clone, Default)]
pub struct Recordings {
state: Arc<Mutex<RecordingSnapshot>>,
}
#[cfg(any(test, feature = "test-utils"))]
impl Recordings {
pub fn snapshot(&self) -> RecordingSnapshot {
self.state.lock().clone()
}
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);
}
}
#[cfg(any(test, feature = "test-utils"))]
#[derive(Clone)]
pub struct RecordingContext<E> {
pub inner: E,
pub recordings: Recordings,
}
#[cfg(any(test, feature = "test-utils"))]
impl<E> RecordingContext<E> {
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<E: Spawner> Spawner for RecordingContext<E> {
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<F, Fut, T>(self, f: F) -> Handle<T>
where
F: FnOnce(Self) -> Fut + Send + 'static,
Fut: Future<Output = T> + 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<std::time::Duration>) -> Result<(), Error> {
self.inner.stop(value, timeout).await
}
fn stopped(&self) -> Signal {
self.inner.stopped()
}
}
#[cfg(any(test, feature = "test-utils"))]
impl<E: Storage> Storage for RecordingContext<E> {
type Blob = RecordingBlob<E::Blob>;
async fn open_versioned(
&self,
partition: &str,
name: &[u8],
versions: std::ops::RangeInclusive<BlobVersion>,
) -> 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<Vec<Vec<u8>>, Error> {
self.inner.scan(partition).await
}
}
#[cfg(any(test, feature = "test-utils"))]
#[derive(Clone)]
pub struct RecordingBlob<B> {
inner: B,
recordings: Recordings,
}
#[cfg(any(test, feature = "test-utils"))]
impl<B: Blob> Blob for RecordingBlob<B> {
async fn read_at_buf(
&self,
offset: u64,
len: usize,
bufs: impl Into<IoBufsMut> + Send,
options: ReadOptions,
) -> Result<IoBufsMut, Error> {
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<IoBufsMut, Error> {
self.recordings.read(options);
self.inner.read_at(offset, len, options).await
}
async fn write_at(
&self,
offset: u64,
bufs: impl Into<IoBufs> + 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
}
}
#[derive(Clone)]
pub struct DelayedSyncContext<E> {
pub inner: E,
pub pending: PendingSyncs,
}
forward_context!(DelayedSyncContext, pending);
impl<E: Spawner> Spawner for DelayedSyncContext<E> {
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<F, Fut, T>(self, f: F) -> Handle<T>
where
F: FnOnce(Self) -> Fut + Send + 'static,
Fut: Future<Output = T> + 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<std::time::Duration>) -> Result<(), Error> {
self.inner.stop(value, timeout).await
}
fn stopped(&self) -> Signal {
self.inner.stopped()
}
}
impl<E: Storage> Storage for DelayedSyncContext<E> {
type Blob = DelayedSyncBlob<E::Blob>;
async fn open_versioned(
&self,
partition: &str,
name: &[u8],
versions: std::ops::RangeInclusive<BlobVersion>,
) -> 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<Vec<Vec<u8>>, Error> {
self.inner.scan(partition).await
}
}
#[derive(Clone)]
pub struct DelayedSyncBlob<B> {
inner: B,
pending: PendingSyncs,
}
impl<B> DelayedSyncBlob<B> {
pub fn new(inner: B) -> (Self, PendingSyncs) {
let pending = PendingSyncs::default();
(
Self {
inner,
pending: pending.clone(),
},
pending,
)
}
}
impl<B: Blob> Blob for DelayedSyncBlob<B> {
async fn read_at_buf(
&self,
offset: u64,
len: usize,
bufs: impl Into<IoBufsMut> + Send,
options: ReadOptions,
) -> Result<IoBufsMut, Error> {
self.inner.read_at_buf(offset, len, bufs, options).await
}
async fn read_at(
&self,
offset: u64,
len: usize,
options: ReadOptions,
) -> Result<IoBufsMut, Error> {
self.inner.read_at(offset, len, options).await
}
async fn write_at(
&self,
offset: u64,
bufs: impl Into<IoBufs> + 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;
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(())
})
}
}
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)
}
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::<Vec<_>>()
};
for sync in syncs {
let _ = sync.release.send(Ok(()));
}
}
pub fn release_pending_syncs(pending: &PendingSyncs) {
for sync in mem::take(&mut *pending.lock()) {
let _ = sync.release.send(Ok(()));
}
}
pub async fn drive_pending_syncs<T>(pending: &PendingSyncs, fut: impl Future<Output = T>) -> 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 => {
release_pending_syncs(pending);
cx.waker().wake_by_ref();
Poll::Pending
}
})
.await
}
pub fn fail_pending_syncs(pending: &PendingSyncs) {
for sync in mem::take(&mut *pending.lock()) {
let _ = sync.release.send(Err(injected_sync_failure()));
}
}
fn injected_sync_failure() -> Error {
Error::Io(std::io::Error::other("injected sync failure").into())
}
struct SyncWaiter {
entered: oneshot::Sender<()>,
release: oneshot::Receiver<Result<(), Error>>,
}
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<SyncWaiter>,
}
impl PendingSyncs {
pub fn lock(&self) -> commonware_utils::sync::MappedMutexGuard<'_, Vec<DeferredSync>> {
commonware_utils::sync::MutexGuard::map(self.state.lock(), |state| &mut state.syncs)
}
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);
}
pub fn calls(&self) -> usize {
self.state.lock().gate.calls
}
fn tracking(&self) -> bool {
self.state.lock().gate.tracking
}
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);
}
}
pub fn arm_fail(&self) {
self.state.lock().fail = true;
}
pub fn starts(&self) -> usize {
self.state.lock().starts
}
pub fn entered(&self) -> usize {
self.state.lock().entered
}
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(()),
}
}
}
#[derive(Clone, Default)]
pub struct WriteFaults {
state: Arc<Mutex<WriteFaultState>>,
}
#[derive(Default)]
struct WriteFaultState {
fail: bool,
writes: u64,
}
impl WriteFaults {
pub fn arm(&self) {
self.state.lock().fail = true;
}
pub fn disarm(&self) {
self.state.lock().fail = false;
}
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;
}
}
#[derive(Clone)]
pub struct WriteFaultContext<E> {
pub inner: E,
pub faults: WriteFaults,
}
forward_context!(WriteFaultContext, faults);
impl<E: Storage> Storage for WriteFaultContext<E> {
type Blob = WriteFaultBlob<E::Blob>;
async fn open_versioned(
&self,
partition: &str,
name: &[u8],
versions: std::ops::RangeInclusive<BlobVersion>,
) -> 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<Vec<Vec<u8>>, Error> {
self.inner.scan(partition).await
}
}
#[derive(Clone)]
pub struct WriteFaultBlob<B> {
inner: B,
faults: WriteFaults,
}
impl<B: Blob> Blob for WriteFaultBlob<B> {
async fn read_at_buf(
&self,
offset: u64,
len: usize,
bufs: impl Into<IoBufsMut> + Send,
options: ReadOptions,
) -> Result<IoBufsMut, Error> {
self.inner.read_at_buf(offset, len, bufs, options).await
}
async fn read_at(
&self,
offset: u64,
len: usize,
options: ReadOptions,
) -> Result<IoBufsMut, Error> {
self.inner.read_at(offset, len, options).await
}
async fn write_at(
&self,
offset: u64,
bufs: impl Into<IoBufs> + 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
}
}
#[derive(Clone)]
pub struct SyncFaultContext<E> {
pub inner: E,
pub fail_partition: String,
}
forward_context!(SyncFaultContext, fail_partition);
impl<E: Storage> Storage for SyncFaultContext<E> {
type Blob = SyncFaultBlob<E::Blob>;
async fn open_versioned(
&self,
partition: &str,
name: &[u8],
versions: std::ops::RangeInclusive<BlobVersion>,
) -> 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<Vec<Vec<u8>>, Error> {
self.inner.scan(partition).await
}
}
#[derive(Clone)]
pub struct SyncFaultBlob<B> {
inner: B,
faulty: bool,
}
impl<B: Blob> Blob for SyncFaultBlob<B> {
async fn read_at_buf(
&self,
offset: u64,
len: usize,
bufs: impl Into<IoBufsMut> + Send,
options: ReadOptions,
) -> Result<IoBufsMut, Error> {
self.inner.read_at_buf(offset, len, bufs, options).await
}
async fn read_at(
&self,
offset: u64,
len: usize,
options: ReadOptions,
) -> Result<IoBufsMut, Error> {
self.inner.read_at(offset, len, options).await
}
async fn write_at(
&self,
offset: u64,
bufs: impl Into<IoBufs> + 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<E: Storage>(
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 {
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);
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 {
assert!(sink.send(b"7 bytes".as_slice()).await.is_ok());
let handle = context.child("recv").spawn(|_| async move {
let _ = stream.recv(5).await;
let _ = stream.recv(5).await;
});
context.sleep(Duration::from_millis(50)).await;
handle.abort();
assert!(matches!(handle.await, Err(Error::Closed)));
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);
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();
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();
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 {
sink.send(b"hello world".as_slice()).await.unwrap();
let received = stream.recv(5).await.unwrap();
assert_eq!(received.coalesce(), b"hello");
assert_eq!(stream.peek(100), b" world");
assert_eq!(stream.peek(3), b" wo");
assert_eq!(stream.peek(100), b" world");
let received = stream.recv(6).await.unwrap();
assert_eq!(received.coalesce(), b" world");
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 {
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
});
context.sleep(Duration::from_millis(10)).await;
sink.send(b"ABCDEFGHIJ".as_slice()).await.unwrap();
let received = recv_handle.await.unwrap();
assert_eq!(received.coalesce(), b"ABC");
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 {
sink.send(b"aaa".as_slice()).await.unwrap();
sink.send(b"bbb".as_slice()).await.unwrap();
sink.send(b"ccc".as_slice()).await.unwrap();
let received = stream.recv(4).await.unwrap();
assert_eq!(received.coalesce(), b"aaab");
assert_eq!(stream.peek(100), b"bbccc");
});
}
#[test]
fn test_buffer_size_limit() {
let (mut sink, mut stream) = Channel::init_with_buffer_size(10);
let executor = deterministic::Runner::default();
executor.start(|context| async move {
let send_handle = context.child("sender").spawn(|_| async move {
sink.send(b"0123456789ABCDEF".as_slice()).await.unwrap();
sink
});
let received = stream.recv(2).await.unwrap();
assert_eq!(received.coalesce(), b"01");
assert_eq!(stream.peek(100), b"23456789");
let received = stream.recv(8).await.unwrap();
assert_eq!(received.coalesce(), b"23456789");
let received = stream.recv(2).await.unwrap();
assert_eq!(received.coalesce(), b"AB");
assert_eq!(stream.peek(100), b"CDEF");
send_handle.await.unwrap();
});
}
#[test]
fn test_recv_before_send() {
let (mut sink, mut stream) = Channel::init_with_buffer_size(10);
let executor = deterministic::Runner::default();
executor.start(|context| async move {
let recv_handle = context
.child("recv")
.spawn(|_| async move { stream.recv(3).await.unwrap() });
context.sleep(Duration::from_millis(10)).await;
sink.send(b"ABCDEFGHIJKLMNOP".as_slice()).await.unwrap();
let received = recv_handle.await.unwrap();
assert_eq!(received.coalesce(), b"ABC");
});
}
}