1#[cfg(any(test, feature = "test-utils"))]
4pub use crate::storage::memory::Storage as MemoryStorage;
5use crate::{
6 Blob, BlobVersion, BufMut, BufferPool, BufferPooler, Clock, Error, Handle, IoBufs, IoBufsMut,
7 Metrics, Name, ReadOptions, Spawner, Storage, Supervisor, WriteOptions,
8 signal::Signal,
9 telemetry::metrics::{Metric, Registered},
10};
11use bytes::{Bytes, BytesMut};
12use commonware_utils::{
13 channel::{fallible::OneshotExt, oneshot},
14 sync::Mutex,
15};
16use governor::clock::{Clock as GovernorClock, ReasonablyRealtime};
17use rand::{TryCryptoRng, TryRng};
18use std::{
19 future::{Future, poll_fn},
20 mem,
21 sync::Arc,
22 task::Poll,
23};
24
25const DEFAULT_BUFFER_SIZE: usize = 64 * 1024;
28
29pub struct Channel {
31 buffer: BytesMut,
33
34 waiter: Option<(usize, oneshot::Sender<Bytes>)>,
38
39 buffer_size: usize,
42
43 drain_waiter: Option<oneshot::Sender<()>>,
46
47 sink_alive: bool,
49
50 stream_alive: bool,
52}
53
54impl Channel {
55 pub fn init() -> (Sink, Stream) {
57 Self::init_with_buffer_size(DEFAULT_BUFFER_SIZE)
58 }
59
60 pub fn init_with_buffer_size(buffer_size: usize) -> (Sink, Stream) {
62 let channel = Arc::new(Mutex::new(Self {
63 buffer: BytesMut::new(),
64 waiter: None,
65 buffer_size,
66 drain_waiter: None,
67 sink_alive: true,
68 stream_alive: true,
69 }));
70 (
71 Sink {
72 channel: channel.clone(),
73 state: SinkState::Open,
74 },
75 Stream {
76 channel,
77 buffer: BytesMut::new(),
78 poisoned: false,
79 },
80 )
81 }
82
83 fn restore_front(&mut self, data: Bytes) {
85 if data.is_empty() {
86 return;
87 }
88
89 let mut restored = BytesMut::with_capacity(data.len() + self.buffer.len());
90 restored.extend_from_slice(&data);
91 restored.extend_from_slice(&self.buffer);
92 self.buffer = restored;
93 }
94
95 fn close_sink(&mut self) {
97 self.sink_alive = false;
98
99 self.waiter.take();
101 }
102}
103
104struct RecvWaiterGuard {
105 channel: Arc<Mutex<Channel>>,
106 active: bool,
107}
108
109impl RecvWaiterGuard {
110 const fn new(channel: Arc<Mutex<Channel>>) -> Self {
111 Self {
112 channel,
113 active: true,
114 }
115 }
116
117 const fn disarm(&mut self) {
118 self.active = false;
119 }
120}
121
122impl Drop for RecvWaiterGuard {
123 fn drop(&mut self) {
124 if !self.active {
125 return;
126 }
127
128 self.channel.lock().waiter.take();
129 }
130}
131
132pub struct Sink {
134 channel: Arc<Mutex<Channel>>,
135 state: SinkState,
136}
137
138enum SinkState {
140 Open,
142 Sending,
144 Closed,
146}
147
148impl Sink {
149 fn close(&mut self) {
150 if matches!(self.state, SinkState::Closed) {
151 return;
152 }
153 self.channel.lock().close_sink();
154 self.state = SinkState::Closed;
155 }
156}
157
158impl crate::Sink for Sink {
159 async fn send(&mut self, bufs: impl Into<IoBufs> + Send) -> Result<(), Error> {
160 match self.state {
161 SinkState::Open => {}
162 SinkState::Sending => {
163 self.close();
164 return Err(Error::Closed);
165 }
166 SinkState::Closed => return Err(Error::Closed),
167 }
168
169 let drain_recv = {
170 let mut channel = self.channel.lock();
171
172 if !channel.stream_alive {
174 channel.close_sink();
175 self.state = SinkState::Closed;
176 return Err(Error::SendFailed);
177 }
178
179 channel.buffer.put(bufs.into());
180
181 if channel
184 .waiter
185 .as_ref()
186 .is_some_and(|(requested, _)| *requested <= channel.buffer.len())
187 {
188 let (requested, os_send) = channel.waiter.take().unwrap();
190 let send_amount = channel.buffer.len().min(requested.max(channel.buffer_size));
191 let data = channel.buffer.split_to(send_amount).freeze();
192
193 if let Err(data) = os_send.send(data) {
196 channel.restore_front(data);
197 if !channel.stream_alive {
198 channel.close_sink();
199 self.state = SinkState::Closed;
200 return Err(Error::SendFailed);
201 }
202 }
203 }
204
205 if channel.buffer.len() > channel.buffer_size {
208 assert!(channel.drain_waiter.is_none());
209 let (os_send, os_recv) = oneshot::channel();
210 channel.drain_waiter = Some(os_send);
211 os_recv
212 } else {
213 return Ok(());
214 }
215 };
216
217 self.state = SinkState::Sending;
220
221 match drain_recv.await {
223 Ok(()) => {
224 self.state = SinkState::Open;
225 Ok(())
226 }
227 Err(_) => {
228 self.close();
229 Err(Error::SendFailed)
230 }
231 }
232 }
233}
234
235impl Drop for Sink {
236 fn drop(&mut self) {
237 self.close();
238 }
239}
240
241pub struct Stream {
243 channel: Arc<Mutex<Channel>>,
244 buffer: BytesMut,
246 poisoned: bool,
247}
248
249impl crate::Stream for Stream {
250 async fn recv(&mut self, len: usize) -> Result<IoBufs, Error> {
251 if self.poisoned {
252 return Err(Error::Closed);
253 }
254
255 let os_recv = {
256 let mut channel = self.channel.lock();
257
258 let target = len.max(channel.buffer_size);
260 let pull_amount = channel
261 .buffer
262 .len()
263 .min(target.saturating_sub(self.buffer.len()));
264 if pull_amount > 0 {
265 let data = channel.buffer.split_to(pull_amount);
266 self.buffer.extend_from_slice(&data);
267
268 if channel.buffer.len() <= channel.buffer_size
270 && let Some(sender) = channel.drain_waiter.take()
271 {
272 sender.send_lossy(());
273 }
274 }
275
276 if self.buffer.len() >= len {
278 return Ok(IoBufs::from(self.buffer.split_to(len).freeze()));
279 }
280
281 if !channel.sink_alive {
283 self.poisoned = true;
284 return Err(Error::RecvFailed);
285 }
286
287 let remaining = len - self.buffer.len();
289 assert!(channel.waiter.is_none());
290 let (os_send, os_recv) = oneshot::channel();
291 channel.waiter = Some((remaining, os_send));
292 os_recv
293 };
294
295 let mut waiter_guard = RecvWaiterGuard::new(self.channel.clone());
296
297 self.poisoned = true;
299
300 let data = match os_recv.await {
302 Ok(data) => {
303 waiter_guard.disarm();
304 self.poisoned = false;
305 data
306 }
307 Err(_) => {
308 waiter_guard.disarm();
309 return Err(Error::RecvFailed);
310 }
311 };
312 self.buffer.extend_from_slice(&data);
313
314 assert!(self.buffer.len() >= len);
315 Ok(IoBufs::from(self.buffer.split_to(len).freeze()))
316 }
317
318 fn peek(&self, max_len: usize) -> &[u8] {
319 let len = max_len.min(self.buffer.len());
320 &self.buffer[..len]
321 }
322}
323
324impl Drop for Stream {
325 fn drop(&mut self) {
326 let mut channel = self.channel.lock();
327 channel.stream_alive = false;
328
329 channel.drain_waiter.take();
331 }
332}
333
334pub struct DeferredSync {
336 pub release: oneshot::Sender<Result<(), Error>>,
338
339 pub blocked: oneshot::Receiver<()>,
341}
342
343#[derive(Clone, Default)]
351pub struct PendingSyncs {
352 state: Arc<Mutex<State>>,
353}
354
355#[derive(Default)]
357struct State {
358 syncs: Vec<DeferredSync>,
360 gate: SyncGateState,
362 unblocked: bool,
364 fail: bool,
366 starts: usize,
368 entered: usize,
370 completions: usize,
372}
373
374impl State {
375 fn defer(&mut self) -> SyncWaiter {
377 let (release, release_rx) = oneshot::channel();
378 let (entered, blocked) = oneshot::channel();
379 self.syncs.push(DeferredSync { release, blocked });
380 SyncWaiter {
381 entered,
382 release: release_rx,
383 }
384 }
385
386 const fn observe(&mut self) -> Option<SyncWaiter> {
389 if !self.gate.tracking {
390 return None;
391 }
392 self.gate.calls += 1;
393 self.gate.waiter.take()
394 }
395
396 fn park(&mut self) -> Option<SyncWaiter> {
398 if self.unblocked {
399 return None;
400 }
401 Some(self.defer())
402 }
403}
404
405macro_rules! forward_context {
410 ($wrapper:ident, $field:ident) => {
411 impl<E: Supervisor> Supervisor for $wrapper<E> {
412 fn name(&self) -> Name {
413 self.inner.name()
414 }
415
416 fn child(&self, label: &'static str) -> Self {
417 Self {
418 inner: self.inner.child(label),
419 $field: self.$field.clone(),
420 }
421 }
422
423 fn with_attribute(self, key: &'static str, value: impl std::fmt::Display) -> Self {
424 Self {
425 inner: self.inner.with_attribute(key, value),
426 $field: self.$field,
427 }
428 }
429 }
430
431 impl<E: Clock> Clock for $wrapper<E> {
432 fn current(&self) -> std::time::SystemTime {
433 self.inner.current()
434 }
435
436 fn sleep(
437 &self,
438 duration: std::time::Duration,
439 ) -> impl Future<Output = ()> + Send + 'static {
440 self.inner.sleep(duration)
441 }
442
443 fn sleep_until(
444 &self,
445 deadline: std::time::SystemTime,
446 ) -> impl Future<Output = ()> + Send + 'static {
447 self.inner.sleep_until(deadline)
448 }
449 }
450
451 impl<E: Clock> GovernorClock for $wrapper<E> {
452 type Instant = std::time::SystemTime;
453
454 fn now(&self) -> Self::Instant {
455 self.current()
456 }
457 }
458
459 impl<E: Clock> ReasonablyRealtime for $wrapper<E> {}
460
461 impl<E: Metrics> Metrics for $wrapper<E> {
462 fn register<N: Into<String>, H: Into<String>, M: Metric>(
463 &self,
464 name: N,
465 help: H,
466 metric: M,
467 ) -> Registered<M> {
468 self.inner.register(name, help, metric)
469 }
470
471 fn encode(&self) -> String {
472 self.inner.encode()
473 }
474 }
475
476 impl<E: BufferPooler> BufferPooler for $wrapper<E> {
477 fn network_buffer_pool(&self) -> &BufferPool {
478 self.inner.network_buffer_pool()
479 }
480
481 fn storage_buffer_pool(&self) -> &BufferPool {
482 self.inner.storage_buffer_pool()
483 }
484 }
485
486 impl<E: TryRng> TryRng for $wrapper<E> {
487 type Error = E::Error;
488
489 fn try_next_u32(&mut self) -> Result<u32, Self::Error> {
490 self.inner.try_next_u32()
491 }
492
493 fn try_next_u64(&mut self) -> Result<u64, Self::Error> {
494 self.inner.try_next_u64()
495 }
496
497 fn try_fill_bytes(&mut self, dest: &mut [u8]) -> Result<(), Self::Error> {
498 self.inner.try_fill_bytes(dest)
499 }
500 }
501
502 impl<E: TryCryptoRng> TryCryptoRng for $wrapper<E> {}
503 };
504}
505
506#[cfg(any(test, feature = "test-utils"))]
508#[derive(Clone, Debug, Default, Eq, PartialEq)]
509pub struct RecordingSnapshot {
510 pub reads: Vec<ReadOptions>,
512 pub writes: Vec<WriteOptions>,
514}
515
516#[cfg(any(test, feature = "test-utils"))]
518#[derive(Clone, Default)]
519pub struct Recordings {
520 state: Arc<Mutex<RecordingSnapshot>>,
521}
522
523#[cfg(any(test, feature = "test-utils"))]
524impl Recordings {
525 pub fn snapshot(&self) -> RecordingSnapshot {
527 self.state.lock().clone()
528 }
529
530 pub fn clear(&self) {
532 *self.state.lock() = RecordingSnapshot::default();
533 }
534
535 fn read(&self, options: ReadOptions) {
536 self.state.lock().reads.push(options);
537 }
538
539 fn write(&self, options: WriteOptions) {
540 self.state.lock().writes.push(options);
541 }
542}
543
544#[cfg(any(test, feature = "test-utils"))]
546#[derive(Clone)]
547pub struct RecordingContext<E> {
548 pub inner: E,
550 pub recordings: Recordings,
552}
553
554#[cfg(any(test, feature = "test-utils"))]
555impl<E> RecordingContext<E> {
556 pub fn new(inner: E) -> (Self, Recordings) {
558 let recordings = Recordings::default();
559 (
560 Self {
561 inner,
562 recordings: recordings.clone(),
563 },
564 recordings,
565 )
566 }
567}
568
569#[cfg(any(test, feature = "test-utils"))]
570forward_context!(RecordingContext, recordings);
571
572#[cfg(any(test, feature = "test-utils"))]
573impl<E: Spawner> Spawner for RecordingContext<E> {
574 fn shared(mut self, blocking: bool) -> Self {
575 self.inner = self.inner.shared(blocking);
576 self
577 }
578
579 fn dedicated(mut self) -> Self {
580 self.inner = self.inner.dedicated();
581 self
582 }
583
584 fn spawn<F, Fut, T>(self, f: F) -> Handle<T>
585 where
586 F: FnOnce(Self) -> Fut + Send + 'static,
587 Fut: Future<Output = T> + Send + 'static,
588 T: Send + 'static,
589 {
590 let recordings = self.recordings;
591 self.inner.spawn(move |inner| f(Self { inner, recordings }))
592 }
593
594 async fn stop(self, value: i32, timeout: Option<std::time::Duration>) -> Result<(), Error> {
595 self.inner.stop(value, timeout).await
596 }
597
598 fn stopped(&self) -> Signal {
599 self.inner.stopped()
600 }
601}
602
603#[cfg(any(test, feature = "test-utils"))]
604impl<E: Storage> Storage for RecordingContext<E> {
605 type Blob = RecordingBlob<E::Blob>;
606
607 async fn open_versioned(
608 &self,
609 partition: &str,
610 name: &[u8],
611 versions: std::ops::RangeInclusive<BlobVersion>,
612 ) -> Result<(Self::Blob, u64, BlobVersion), Error> {
613 let (inner, len, version) = self.inner.open_versioned(partition, name, versions).await?;
614 Ok((
615 RecordingBlob {
616 inner,
617 recordings: self.recordings.clone(),
618 },
619 len,
620 version,
621 ))
622 }
623
624 async fn remove(&self, partition: &str, name: Option<&[u8]>) -> Result<(), Error> {
625 self.inner.remove(partition, name).await
626 }
627
628 async fn scan(&self, partition: &str) -> Result<Vec<Vec<u8>>, Error> {
629 self.inner.scan(partition).await
630 }
631}
632
633#[cfg(any(test, feature = "test-utils"))]
635#[derive(Clone)]
636pub struct RecordingBlob<B> {
637 inner: B,
638 recordings: Recordings,
639}
640
641#[cfg(any(test, feature = "test-utils"))]
642impl<B: Blob> Blob for RecordingBlob<B> {
643 async fn read_at_buf(
644 &self,
645 offset: u64,
646 len: usize,
647 bufs: impl Into<IoBufsMut> + Send,
648 options: ReadOptions,
649 ) -> Result<IoBufsMut, Error> {
650 self.recordings.read(options);
651 self.inner.read_at_buf(offset, len, bufs, options).await
652 }
653
654 async fn read_at(
655 &self,
656 offset: u64,
657 len: usize,
658 options: ReadOptions,
659 ) -> Result<IoBufsMut, Error> {
660 self.recordings.read(options);
661 self.inner.read_at(offset, len, options).await
662 }
663
664 async fn write_at(
665 &self,
666 offset: u64,
667 bufs: impl Into<IoBufs> + Send,
668 options: WriteOptions,
669 ) -> Result<(), Error> {
670 self.recordings.write(options);
671 self.inner.write_at(offset, bufs, options).await
672 }
673
674 async fn resize(&self, len: u64) -> Result<(), Error> {
675 self.inner.resize(len).await
676 }
677
678 async fn sync(&self) -> Result<(), Error> {
679 self.inner.sync().await
680 }
681
682 async fn start_sync(&self) -> Handle<()> {
683 self.inner.start_sync().await
684 }
685}
686
687#[derive(Clone)]
689pub struct DelayedSyncContext<E> {
690 pub inner: E,
691 pub pending: PendingSyncs,
692}
693
694forward_context!(DelayedSyncContext, pending);
695
696impl<E: Spawner> Spawner for DelayedSyncContext<E> {
697 fn shared(mut self, blocking: bool) -> Self {
698 self.inner = self.inner.shared(blocking);
699 self
700 }
701
702 fn dedicated(mut self) -> Self {
703 self.inner = self.inner.dedicated();
704 self
705 }
706
707 fn spawn<F, Fut, T>(self, f: F) -> Handle<T>
708 where
709 F: FnOnce(Self) -> Fut + Send + 'static,
710 Fut: Future<Output = T> + Send + 'static,
711 T: Send + 'static,
712 {
713 let pending = self.pending;
714 self.inner.spawn(move |inner| f(Self { inner, pending }))
715 }
716
717 async fn stop(self, value: i32, timeout: Option<std::time::Duration>) -> Result<(), Error> {
718 self.inner.stop(value, timeout).await
719 }
720
721 fn stopped(&self) -> Signal {
722 self.inner.stopped()
723 }
724}
725
726impl<E: Storage> Storage for DelayedSyncContext<E> {
727 type Blob = DelayedSyncBlob<E::Blob>;
728
729 async fn open_versioned(
730 &self,
731 partition: &str,
732 name: &[u8],
733 versions: std::ops::RangeInclusive<BlobVersion>,
734 ) -> Result<(Self::Blob, u64, BlobVersion), Error> {
735 let (inner, len, version) = self.inner.open_versioned(partition, name, versions).await?;
736 Ok((
737 DelayedSyncBlob {
738 inner,
739 pending: self.pending.clone(),
740 },
741 len,
742 version,
743 ))
744 }
745
746 async fn remove(&self, partition: &str, name: Option<&[u8]>) -> Result<(), Error> {
747 self.inner.remove(partition, name).await
748 }
749
750 async fn scan(&self, partition: &str) -> Result<Vec<Vec<u8>>, Error> {
751 self.inner.scan(partition).await
752 }
753}
754
755#[derive(Clone)]
757pub struct DelayedSyncBlob<B> {
758 inner: B,
759 pending: PendingSyncs,
760}
761
762impl<B> DelayedSyncBlob<B> {
763 pub fn new(inner: B) -> (Self, PendingSyncs) {
765 let pending = PendingSyncs::default();
766 (
767 Self {
768 inner,
769 pending: pending.clone(),
770 },
771 pending,
772 )
773 }
774}
775
776impl<B: Blob> Blob for DelayedSyncBlob<B> {
777 async fn read_at_buf(
778 &self,
779 offset: u64,
780 len: usize,
781 bufs: impl Into<IoBufsMut> + Send,
782 options: ReadOptions,
783 ) -> Result<IoBufsMut, Error> {
784 self.inner.read_at_buf(offset, len, bufs, options).await
785 }
786
787 async fn read_at(
788 &self,
789 offset: u64,
790 len: usize,
791 options: ReadOptions,
792 ) -> Result<IoBufsMut, Error> {
793 self.inner.read_at(offset, len, options).await
794 }
795
796 async fn write_at(
797 &self,
798 offset: u64,
799 bufs: impl Into<IoBufs> + Send,
800 options: WriteOptions,
801 ) -> Result<(), Error> {
802 if !options.contains(WriteOptions::SYNC) || !self.pending.tracking() {
803 return self.inner.write_at(offset, bufs, options).await;
804 }
805 self.inner
806 .write_at(offset, bufs, options.without(WriteOptions::SYNC))
807 .await?;
808 self.sync().await
809 }
810
811 async fn resize(&self, len: u64) -> Result<(), Error> {
812 self.inner.resize(len).await
813 }
814
815 async fn sync(&self) -> Result<(), Error> {
816 self.pending.wait().await?;
817 self.inner.sync().await
818 }
819
820 async fn start_sync(&self) -> Handle<()> {
821 let pending = self.pending.clone();
822 let inner = self.inner.clone();
823 let waiter = {
824 let mut state = pending.state.lock();
825 state.starts += 1;
826 state.observe().or_else(|| state.park())
828 };
829 Handle::from_future(async move {
830 let fail = {
831 let mut state = pending.state.lock();
832 state.entered += 1;
833 state.fail
834 };
835 match waiter {
836 Some(waiter) => waiter.wait().await?,
837 None if fail => return Err(injected_sync_failure()),
838 None => {}
839 }
840 inner.sync().await?;
841 pending.state.lock().completions += 1;
842 Ok(())
843 })
844 }
845}
846
847pub fn next_pending_sync(pending: &PendingSyncs) -> DeferredSync {
849 let mut pending = pending.lock();
850 assert!(!pending.is_empty(), "no pending sync was started");
851 pending.remove(0)
852}
853
854pub fn release_next_pending_syncs(pending: &PendingSyncs, count: usize) {
856 let syncs = {
857 let mut pending = pending.lock();
858 assert!(
859 pending.len() >= count,
860 "not enough pending syncs: have {}, need {count}",
861 pending.len()
862 );
863 pending.drain(..count).collect::<Vec<_>>()
864 };
865 for sync in syncs {
866 let _ = sync.release.send(Ok(()));
867 }
868}
869
870pub fn release_pending_syncs(pending: &PendingSyncs) {
872 for sync in mem::take(&mut *pending.lock()) {
873 let _ = sync.release.send(Ok(()));
874 }
875}
876
877pub async fn drive_pending_syncs<T>(pending: &PendingSyncs, fut: impl Future<Output = T>) -> T {
879 let mut fut = std::pin::pin!(fut);
880 poll_fn(|cx| match fut.as_mut().poll(cx) {
881 Poll::Ready(out) => Poll::Ready(out),
882 Poll::Pending => {
883 release_pending_syncs(pending);
886 cx.waker().wake_by_ref();
887 Poll::Pending
888 }
889 })
890 .await
891}
892
893pub fn fail_pending_syncs(pending: &PendingSyncs) {
895 for sync in mem::take(&mut *pending.lock()) {
896 let _ = sync.release.send(Err(injected_sync_failure()));
897 }
898}
899
900fn injected_sync_failure() -> Error {
902 Error::Io(std::io::Error::other("injected sync failure").into())
903}
904
905struct SyncWaiter {
906 entered: oneshot::Sender<()>,
907 release: oneshot::Receiver<Result<(), Error>>,
908}
909
910impl SyncWaiter {
911 async fn wait(self) -> Result<(), Error> {
912 self.entered.send_lossy(());
913 self.release.await.map_err(|_| Error::Closed)??;
914 Ok(())
915 }
916}
917
918#[derive(Default)]
919struct SyncGateState {
920 tracking: bool,
921 calls: usize,
922 waiter: Option<SyncWaiter>,
923}
924
925impl PendingSyncs {
926 pub fn lock(&self) -> commonware_utils::sync::MappedMutexGuard<'_, Vec<DeferredSync>> {
928 commonware_utils::sync::MutexGuard::map(self.state.lock(), |state| &mut state.syncs)
929 }
930
931 pub fn arm(&self) {
937 let mut state = self.state.lock();
938 assert!(!state.gate.tracking, "sync gate already armed");
939 assert!(
940 state.gate.waiter.is_none(),
941 "sync gate already has a waiter"
942 );
943 state.gate.tracking = true;
944 state.gate.calls = 0;
945 let waiter = state.defer();
946 state.gate.waiter = Some(waiter);
947 }
948
949 pub fn calls(&self) -> usize {
951 self.state.lock().gate.calls
952 }
953
954 fn tracking(&self) -> bool {
955 self.state.lock().gate.tracking
956 }
957
958 pub fn unblock(&self) {
963 let (drained, fail) = {
964 let mut state = self.state.lock();
965 state.unblocked = true;
966 (mem::take(&mut state.syncs), state.fail)
967 };
968 for sync in drained {
969 let result = if fail {
970 Err(injected_sync_failure())
971 } else {
972 Ok(())
973 };
974 let _ = sync.release.send(result);
975 }
976 }
977
978 pub fn arm_fail(&self) {
981 self.state.lock().fail = true;
982 }
983
984 pub fn starts(&self) -> usize {
986 self.state.lock().starts
987 }
988
989 pub fn entered(&self) -> usize {
991 self.state.lock().entered
992 }
993
994 pub fn completions(&self) -> usize {
996 self.state.lock().completions
997 }
998
999 async fn wait(&self) -> Result<(), Error> {
1000 let waiter = self.state.lock().observe();
1001 match waiter {
1002 Some(waiter) => waiter.wait().await,
1003 None => Ok(()),
1004 }
1005 }
1006}
1007
1008#[derive(Clone, Default)]
1011pub struct WriteFaults {
1012 state: Arc<Mutex<WriteFaultState>>,
1013}
1014
1015#[derive(Default)]
1016struct WriteFaultState {
1017 fail: bool,
1018 writes: u64,
1019}
1020
1021impl WriteFaults {
1022 pub fn arm(&self) {
1024 self.state.lock().fail = true;
1025 }
1026
1027 pub fn disarm(&self) {
1029 self.state.lock().fail = false;
1030 }
1031
1032 pub fn writes(&self) -> u64 {
1034 self.state.lock().writes
1035 }
1036
1037 fn check(&self) -> Result<(), Error> {
1038 if self.state.lock().fail {
1039 return Err(Error::Io(
1040 std::io::Error::other("injected write failure").into(),
1041 ));
1042 }
1043 Ok(())
1044 }
1045
1046 fn note(&self) {
1047 self.state.lock().writes += 1;
1048 }
1049}
1050
1051#[derive(Clone)]
1055pub struct WriteFaultContext<E> {
1056 pub inner: E,
1057 pub faults: WriteFaults,
1058}
1059
1060forward_context!(WriteFaultContext, faults);
1061
1062impl<E: Storage> Storage for WriteFaultContext<E> {
1063 type Blob = WriteFaultBlob<E::Blob>;
1064
1065 async fn open_versioned(
1066 &self,
1067 partition: &str,
1068 name: &[u8],
1069 versions: std::ops::RangeInclusive<BlobVersion>,
1070 ) -> Result<(Self::Blob, u64, BlobVersion), Error> {
1071 let (inner, len, version) = self.inner.open_versioned(partition, name, versions).await?;
1072 Ok((
1073 WriteFaultBlob {
1074 inner,
1075 faults: self.faults.clone(),
1076 },
1077 len,
1078 version,
1079 ))
1080 }
1081
1082 async fn remove(&self, partition: &str, name: Option<&[u8]>) -> Result<(), Error> {
1083 self.inner.remove(partition, name).await
1084 }
1085
1086 async fn scan(&self, partition: &str) -> Result<Vec<Vec<u8>>, Error> {
1087 self.inner.scan(partition).await
1088 }
1089}
1090
1091#[derive(Clone)]
1093pub struct WriteFaultBlob<B> {
1094 inner: B,
1095 faults: WriteFaults,
1096}
1097
1098impl<B: Blob> Blob for WriteFaultBlob<B> {
1099 async fn read_at_buf(
1100 &self,
1101 offset: u64,
1102 len: usize,
1103 bufs: impl Into<IoBufsMut> + Send,
1104 options: ReadOptions,
1105 ) -> Result<IoBufsMut, Error> {
1106 self.inner.read_at_buf(offset, len, bufs, options).await
1107 }
1108
1109 async fn read_at(
1110 &self,
1111 offset: u64,
1112 len: usize,
1113 options: ReadOptions,
1114 ) -> Result<IoBufsMut, Error> {
1115 self.inner.read_at(offset, len, options).await
1116 }
1117
1118 async fn write_at(
1119 &self,
1120 offset: u64,
1121 bufs: impl Into<IoBufs> + Send,
1122 options: WriteOptions,
1123 ) -> Result<(), Error> {
1124 self.faults.check()?;
1125 self.inner.write_at(offset, bufs, options).await?;
1126 self.faults.note();
1127 Ok(())
1128 }
1129
1130 async fn resize(&self, len: u64) -> Result<(), Error> {
1131 self.inner.resize(len).await
1132 }
1133
1134 async fn sync(&self) -> Result<(), Error> {
1135 self.inner.sync().await
1136 }
1137
1138 async fn start_sync(&self) -> Handle<()> {
1139 self.inner.start_sync().await
1140 }
1141}
1142
1143#[derive(Clone)]
1145pub struct SyncFaultContext<E> {
1146 pub inner: E,
1147 pub fail_partition: String,
1148}
1149
1150forward_context!(SyncFaultContext, fail_partition);
1151
1152impl<E: Storage> Storage for SyncFaultContext<E> {
1153 type Blob = SyncFaultBlob<E::Blob>;
1154
1155 async fn open_versioned(
1156 &self,
1157 partition: &str,
1158 name: &[u8],
1159 versions: std::ops::RangeInclusive<BlobVersion>,
1160 ) -> Result<(Self::Blob, u64, BlobVersion), Error> {
1161 let (inner, len, version) = self.inner.open_versioned(partition, name, versions).await?;
1162 Ok((
1163 SyncFaultBlob {
1164 inner,
1165 faulty: partition == self.fail_partition,
1166 },
1167 len,
1168 version,
1169 ))
1170 }
1171
1172 async fn remove(&self, partition: &str, name: Option<&[u8]>) -> Result<(), Error> {
1173 self.inner.remove(partition, name).await
1174 }
1175
1176 async fn scan(&self, partition: &str) -> Result<Vec<Vec<u8>>, Error> {
1177 self.inner.scan(partition).await
1178 }
1179}
1180
1181#[derive(Clone)]
1183pub struct SyncFaultBlob<B> {
1184 inner: B,
1185 faulty: bool,
1186}
1187
1188impl<B: Blob> Blob for SyncFaultBlob<B> {
1189 async fn read_at_buf(
1190 &self,
1191 offset: u64,
1192 len: usize,
1193 bufs: impl Into<IoBufsMut> + Send,
1194 options: ReadOptions,
1195 ) -> Result<IoBufsMut, Error> {
1196 self.inner.read_at_buf(offset, len, bufs, options).await
1197 }
1198
1199 async fn read_at(
1200 &self,
1201 offset: u64,
1202 len: usize,
1203 options: ReadOptions,
1204 ) -> Result<IoBufsMut, Error> {
1205 self.inner.read_at(offset, len, options).await
1206 }
1207
1208 async fn write_at(
1209 &self,
1210 offset: u64,
1211 bufs: impl Into<IoBufs> + Send,
1212 options: WriteOptions,
1213 ) -> Result<(), Error> {
1214 self.inner.write_at(offset, bufs, options).await
1215 }
1216
1217 async fn resize(&self, len: u64) -> Result<(), Error> {
1218 self.inner.resize(len).await
1219 }
1220
1221 async fn sync(&self) -> Result<(), Error> {
1222 if self.faulty {
1223 let err = std::io::Error::other("injected partition sync fault");
1224 return Err(Error::Io(err.into()));
1225 }
1226 self.inner.sync().await
1227 }
1228
1229 async fn start_sync(&self) -> Handle<()> {
1230 if self.faulty {
1231 return Handle::ready(self.sync().await);
1232 }
1233 self.inner.start_sync().await
1234 }
1235}
1236
1237#[cfg(test)]
1238mod tests {
1239 use super::*;
1240 use crate::{Clock, IoBufMut, Runner, Sink, Spawner, Stream, deterministic};
1241 use commonware_macros::select;
1242 use std::{thread::sleep, time::Duration};
1243
1244 #[test]
1245 fn recording_context_preserves_data_and_records_options() {
1246 deterministic::Runner::default().start(|context| async move {
1247 let (context, recordings) = RecordingContext::new(context);
1248 let (blob, _) = context.open("recording", b"blob").await.unwrap();
1249
1250 blob.write_at(0, b"data", WriteOptions::DONT_CACHE)
1251 .await
1252 .unwrap();
1253 let read = blob.read_at(0, 4, ReadOptions::DONT_CACHE).await.unwrap();
1254 assert_eq!(read.coalesce(), b"data");
1255
1256 let read = blob
1257 .read_at_buf(0, 4, IoBufMut::with_capacity(4), ReadOptions::default())
1258 .await
1259 .unwrap();
1260 assert_eq!(read.coalesce(), b"data");
1261
1262 assert_eq!(
1263 recordings.snapshot(),
1264 RecordingSnapshot {
1265 reads: vec![ReadOptions::DONT_CACHE, ReadOptions::default()],
1266 writes: vec![WriteOptions::DONT_CACHE],
1267 }
1268 );
1269 recordings.clear();
1270 assert_eq!(recordings.snapshot(), RecordingSnapshot::default());
1271 });
1272 }
1273
1274 async fn assert_read_options_forwarded<E: Storage>(
1275 context: &E,
1276 recordings: &Recordings,
1277 partition: &str,
1278 ) {
1279 let (blob, _) = context.open(partition, b"blob").await.unwrap();
1280 blob.write_at(0, b"data", WriteOptions::default())
1281 .await
1282 .unwrap();
1283 recordings.clear();
1284
1285 let read = blob.read_at(0, 4, ReadOptions::DONT_CACHE).await.unwrap();
1286 assert_eq!(read.coalesce(), b"data");
1287
1288 let read = blob
1289 .read_at_buf(0, 4, IoBufMut::with_capacity(4), ReadOptions::DONT_CACHE)
1290 .await
1291 .unwrap();
1292 assert_eq!(read.coalesce(), b"data");
1293
1294 assert_eq!(
1295 recordings.snapshot(),
1296 RecordingSnapshot {
1297 reads: vec![ReadOptions::DONT_CACHE, ReadOptions::DONT_CACHE],
1298 writes: Vec::new(),
1299 }
1300 );
1301 }
1302
1303 #[test]
1304 fn delayed_sync_blob_forwards_read_options() {
1305 deterministic::Runner::default().start(|context| async move {
1306 let (inner, recordings) = RecordingContext::new(context);
1307 let context = DelayedSyncContext {
1308 inner,
1309 pending: PendingSyncs::default(),
1310 };
1311
1312 assert_read_options_forwarded(&context, &recordings, "delayed_sync").await;
1313 });
1314 }
1315
1316 #[test]
1317 fn write_fault_blob_forwards_read_options() {
1318 deterministic::Runner::default().start(|context| async move {
1319 let (inner, recordings) = RecordingContext::new(context);
1320 let context = WriteFaultContext {
1321 inner,
1322 faults: WriteFaults::default(),
1323 };
1324
1325 assert_read_options_forwarded(&context, &recordings, "write_fault").await;
1326 });
1327 }
1328
1329 #[test]
1330 fn sync_fault_blob_forwards_read_options() {
1331 deterministic::Runner::default().start(|context| async move {
1332 let (inner, recordings) = RecordingContext::new(context);
1333 let context = SyncFaultContext {
1334 inner,
1335 fail_partition: "sync_fault".to_string(),
1336 };
1337
1338 assert_read_options_forwarded(&context, &recordings, "sync_fault").await;
1339 });
1340 }
1341
1342 #[test]
1343 fn test_send_recv() {
1344 let (mut sink, mut stream) = Channel::init();
1345 let data = b"hello world";
1346
1347 let executor = deterministic::Runner::default();
1348 executor.start(|_| async move {
1349 sink.send(data.as_slice()).await.unwrap();
1350 let received = stream.recv(data.len()).await.unwrap();
1351 assert_eq!(received.coalesce(), data);
1352 });
1353 }
1354
1355 #[test]
1356 fn test_send_recv_partial_multiple() {
1357 let (mut sink, mut stream) = Channel::init();
1358 let data = b"hello";
1359 let data2 = b" world";
1360
1361 let executor = deterministic::Runner::default();
1362 executor.start(|_| async move {
1363 sink.send(data.as_slice()).await.unwrap();
1364 sink.send(data2.as_slice()).await.unwrap();
1365 let received = stream.recv(5).await.unwrap();
1366 assert_eq!(received.coalesce(), b"hello");
1367 let received = stream.recv(5).await.unwrap();
1368 assert_eq!(received.coalesce(), b" worl");
1369 let received = stream.recv(1).await.unwrap();
1370 assert_eq!(received.coalesce(), b"d");
1371 });
1372 }
1373
1374 #[test]
1375 fn test_send_recv_async() {
1376 let (mut sink, mut stream) = Channel::init();
1377 let data = b"hello world";
1378
1379 let executor = deterministic::Runner::default();
1380 executor.start(|_| async move {
1381 let (received, _) = futures::try_join!(stream.recv(data.len()), async {
1382 sleep(Duration::from_millis(50));
1383 sink.send(data.as_slice()).await
1384 })
1385 .unwrap();
1386 assert_eq!(received.coalesce(), data);
1387 });
1388 }
1389
1390 #[test]
1391 fn test_recv_error_sink_dropped_while_waiting() {
1392 let (sink, mut stream) = Channel::init();
1393
1394 let executor = deterministic::Runner::default();
1395 executor.start(|context| async move {
1396 futures::join!(
1397 async {
1398 let result = stream.recv(5).await;
1399 assert!(matches!(result, Err(Error::RecvFailed)));
1400 let result = stream.recv(5).await;
1401 assert!(matches!(result, Err(Error::Closed)));
1402 },
1403 async {
1404 context.sleep(Duration::from_millis(50)).await;
1406 drop(sink);
1407 }
1408 );
1409 });
1410 }
1411
1412 #[test]
1413 fn test_recv_error_sink_dropped_before_recv() {
1414 let (sink, mut stream) = Channel::init();
1415 drop(sink); let executor = deterministic::Runner::default();
1418 executor.start(|_| async move {
1419 let result = stream.recv(5).await;
1420 assert!(matches!(result, Err(Error::RecvFailed)));
1421 let result = stream.recv(5).await;
1422 assert!(matches!(result, Err(Error::Closed)));
1423 });
1424 }
1425
1426 #[test]
1427 fn test_send_error_stream_dropped() {
1428 let (mut sink, mut stream) = Channel::init();
1429
1430 let executor = deterministic::Runner::default();
1431 executor.start(|context| async move {
1432 assert!(sink.send(b"7 bytes".as_slice()).await.is_ok());
1434
1435 let handle = context.child("recv").spawn(|_| async move {
1437 let _ = stream.recv(5).await;
1438 let _ = stream.recv(5).await;
1439 });
1440
1441 context.sleep(Duration::from_millis(50)).await;
1443
1444 handle.abort();
1446 assert!(matches!(handle.await, Err(Error::Closed)));
1447
1448 let result = sink.send(b"hello world".as_slice()).await;
1450 assert!(matches!(result, Err(Error::SendFailed)));
1451 let result = sink.send(b"hello world".as_slice()).await;
1452 assert!(matches!(result, Err(Error::Closed)));
1453 });
1454 }
1455
1456 #[test]
1457 fn test_send_error_stream_dropped_before_send() {
1458 let (mut sink, stream) = Channel::init();
1459 drop(stream); let executor = deterministic::Runner::default();
1462 executor.start(|_| async move {
1463 let result = sink.send(b"hello world".as_slice()).await;
1464 assert!(matches!(result, Err(Error::SendFailed)));
1465 let result = sink.send(b"hello world".as_slice()).await;
1466 assert!(matches!(result, Err(Error::Closed)));
1467 });
1468 }
1469
1470 #[test]
1471 fn test_recv_timeout() {
1472 let (_sink, mut stream) = Channel::init();
1473
1474 let executor = deterministic::Runner::default();
1477 executor.start(|context| async move {
1478 select! {
1479 v = stream.recv(5) => {
1480 panic!("unexpected value: {v:?}");
1481 },
1482 _ = context.sleep(Duration::from_millis(100)) => "timeout",
1483 };
1484 });
1485 }
1486
1487 #[test]
1488 fn test_peek_empty() {
1489 let (_sink, stream) = Channel::init();
1490
1491 assert!(stream.peek(10).is_empty());
1493 }
1494
1495 #[test]
1496 fn test_peek_after_partial_recv() {
1497 let (mut sink, mut stream) = Channel::init();
1498
1499 let executor = deterministic::Runner::default();
1500 executor.start(|_| async move {
1501 sink.send(b"hello world".as_slice()).await.unwrap();
1503
1504 let received = stream.recv(5).await.unwrap();
1506 assert_eq!(received.coalesce(), b"hello");
1507
1508 assert_eq!(stream.peek(100), b" world");
1510
1511 assert_eq!(stream.peek(3), b" wo");
1513
1514 assert_eq!(stream.peek(100), b" world");
1516
1517 let received = stream.recv(6).await.unwrap();
1519 assert_eq!(received.coalesce(), b" world");
1520
1521 assert!(stream.peek(100).is_empty());
1523 });
1524 }
1525
1526 #[test]
1527 fn test_peek_after_recv_wakeup() {
1528 let (mut sink, mut stream) = Channel::init_with_buffer_size(64);
1529
1530 let executor = deterministic::Runner::default();
1531 executor.start(|context| async move {
1532 let (tx, rx) = oneshot::channel();
1534 let recv_handle = context.child("recv").spawn(|_| async move {
1535 let data = stream.recv(3).await.unwrap();
1536 tx.send(stream).ok();
1537 data
1538 });
1539
1540 context.sleep(Duration::from_millis(10)).await;
1542
1543 sink.send(b"ABCDEFGHIJ".as_slice()).await.unwrap();
1545
1546 let received = recv_handle.await.unwrap();
1548 assert_eq!(received.coalesce(), b"ABC");
1549
1550 let stream = rx.await.unwrap();
1552 assert_eq!(stream.peek(100), b"DEFGHIJ");
1553 });
1554 }
1555
1556 #[test]
1557 fn test_peek_multiple_sends() {
1558 let (mut sink, mut stream) = Channel::init();
1559
1560 let executor = deterministic::Runner::default();
1561 executor.start(|_| async move {
1562 sink.send(b"aaa".as_slice()).await.unwrap();
1564 sink.send(b"bbb".as_slice()).await.unwrap();
1565 sink.send(b"ccc".as_slice()).await.unwrap();
1566
1567 let received = stream.recv(4).await.unwrap();
1569 assert_eq!(received.coalesce(), b"aaab");
1570
1571 assert_eq!(stream.peek(100), b"bbccc");
1573 });
1574 }
1575
1576 #[test]
1577 fn test_buffer_size_limit() {
1578 let (mut sink, mut stream) = Channel::init_with_buffer_size(10);
1580
1581 let executor = deterministic::Runner::default();
1582 executor.start(|context| async move {
1583 let send_handle = context.child("sender").spawn(|_| async move {
1586 sink.send(b"0123456789ABCDEF".as_slice()).await.unwrap();
1587 sink
1588 });
1589
1590 let received = stream.recv(2).await.unwrap();
1592 assert_eq!(received.coalesce(), b"01");
1593
1594 assert_eq!(stream.peek(100), b"23456789");
1596
1597 let received = stream.recv(8).await.unwrap();
1600 assert_eq!(received.coalesce(), b"23456789");
1601
1602 let received = stream.recv(2).await.unwrap();
1604 assert_eq!(received.coalesce(), b"AB");
1605
1606 assert_eq!(stream.peek(100), b"CDEF");
1607
1608 send_handle.await.unwrap();
1610 });
1611 }
1612
1613 #[test]
1614 fn test_recv_before_send() {
1615 let (mut sink, mut stream) = Channel::init_with_buffer_size(10);
1617
1618 let executor = deterministic::Runner::default();
1619 executor.start(|context| async move {
1620 let recv_handle = context
1622 .child("recv")
1623 .spawn(|_| async move { stream.recv(3).await.unwrap() });
1624
1625 context.sleep(Duration::from_millis(10)).await;
1627
1628 sink.send(b"ABCDEFGHIJKLMNOP".as_slice()).await.unwrap();
1630
1631 let received = recv_handle.await.unwrap();
1633 assert_eq!(received.coalesce(), b"ABC");
1634 });
1635 }
1636}