1#![cfg(any(test, feature = "testing"))]
14
15use std::{ops::RangeBounds, sync::Arc};
16
17use async_lock::Mutex;
18use async_trait::async_trait;
19use futures::future::Future;
20use hotshot_types::{
21 data::VidShare, simple_certificate::CertificatePair, traits::node_implementation::NodeType,
22};
23
24use super::{
25 Aggregate, AggregatesStorage, AvailabilityStorage, NodeStorage, UpdateAggregatesStorage,
26 UpdateAvailabilityStorage,
27 pruning::{PruneStorage, PrunedHeightStorage, PrunerCfg, PrunerConfig},
28};
29use crate::{
30 Header, Payload, QueryError, QueryResult,
31 availability::{
32 BlockId, BlockQueryData, Certificate2, LeafId, LeafQueryData, NamespaceId,
33 PayloadQueryData, QueryableHeader, QueryablePayload, TransactionHash, VidCommonQueryData,
34 },
35 data_source::{
36 VersionedDataSource,
37 storage::{PayloadMetadata, VidCommonMetadata},
38 update,
39 },
40 metrics::PrometheusMetrics,
41 node::{SyncStatusQueryData, TimeWindowQueryData, WindowStart},
42 status::HasMetrics,
43};
44
45#[derive(Clone, Copy, Debug, PartialEq, Eq)]
47pub enum FailableAction {
48 GetHeader,
51 GetLeaf,
52 GetBlock,
53 GetPayload,
54 GetPayloadMetadata,
55 GetVidCommon,
56 GetVidCommonMetadata,
57 GetHeaderRange,
58 GetLeafRange,
59 GetBlockRange,
60 GetPayloadRange,
61 GetPayloadMetadataRange,
62 GetVidCommonRange,
63 GetVidCommonMetadataRange,
64 GetTransaction,
65 FirstAvailableLeaf,
66 GetStateCert,
67
68 Any,
70}
71
72impl FailableAction {
73 fn matches(self, action: Self) -> bool {
75 self == action || self == Self::Any
78 }
79}
80
81#[derive(Clone, Copy, Debug, Default)]
82enum FailureMode {
83 #[default]
84 Never,
85 Once(FailableAction),
86 Always(FailableAction),
87}
88
89impl FailureMode {
90 fn maybe_fail(&mut self, action: FailableAction) -> QueryResult<()> {
91 match self {
92 Self::Once(fail_action) if fail_action.matches(action) => {
93 *self = Self::Never;
94 },
95 Self::Always(fail_action) if fail_action.matches(action) => {},
96 _ => return Ok(()),
97 }
98
99 Err(QueryError::Error {
100 message: "injected error".into(),
101 })
102 }
103}
104
105#[derive(Debug, Default)]
106struct Failure {
107 on_read: FailureMode,
108 on_write: FailureMode,
109 on_commit: FailureMode,
110 on_begin_writable: FailureMode,
111 on_begin_read_only: FailureMode,
112}
113
114#[derive(Clone, Debug)]
116pub struct FailStorage<S> {
117 inner: S,
118 failure: Arc<Mutex<Failure>>,
119}
120
121impl<S> From<S> for FailStorage<S> {
122 fn from(inner: S) -> Self {
123 Self {
124 inner,
125 failure: Default::default(),
126 }
127 }
128}
129
130impl<S> FailStorage<S> {
131 pub async fn fail_reads(&self, action: FailableAction) {
132 self.failure.lock().await.on_read = FailureMode::Always(action);
133 }
134
135 pub async fn fail_writes(&self, action: FailableAction) {
136 self.failure.lock().await.on_write = FailureMode::Always(action);
137 }
138
139 pub async fn fail_commits(&self, action: FailableAction) {
140 self.failure.lock().await.on_commit = FailureMode::Always(action);
141 }
142
143 pub async fn fail_begins_writable(&self, action: FailableAction) {
144 self.failure.lock().await.on_begin_writable = FailureMode::Always(action);
145 }
146
147 pub async fn fail_begins_read_only(&self, action: FailableAction) {
148 self.failure.lock().await.on_begin_read_only = FailureMode::Always(action);
149 }
150
151 pub async fn fail(&self, action: FailableAction) {
152 let mut failure = self.failure.lock().await;
153 failure.on_read = FailureMode::Always(action);
154 failure.on_write = FailureMode::Always(action);
155 failure.on_commit = FailureMode::Always(action);
156 failure.on_begin_writable = FailureMode::Always(action);
157 failure.on_begin_read_only = FailureMode::Always(action);
158 }
159
160 pub async fn pass_reads(&self) {
161 self.failure.lock().await.on_read = FailureMode::Never;
162 }
163
164 pub async fn pass_writes(&self) {
165 self.failure.lock().await.on_write = FailureMode::Never;
166 }
167
168 pub async fn pass_commits(&self) {
169 self.failure.lock().await.on_commit = FailureMode::Never;
170 }
171
172 pub async fn pass_begins_writable(&self) {
173 self.failure.lock().await.on_begin_writable = FailureMode::Never;
174 }
175
176 pub async fn pass_begins_read_only(&self) {
177 self.failure.lock().await.on_begin_read_only = FailureMode::Never;
178 }
179
180 pub async fn pass(&self) {
181 let mut failure = self.failure.lock().await;
182 failure.on_read = FailureMode::Never;
183 failure.on_write = FailureMode::Never;
184 failure.on_commit = FailureMode::Never;
185 failure.on_begin_writable = FailureMode::Never;
186 failure.on_begin_read_only = FailureMode::Never;
187 }
188
189 pub async fn fail_one_read(&self, action: FailableAction) {
190 self.failure.lock().await.on_read = FailureMode::Once(action);
191 }
192
193 pub async fn fail_one_write(&self, action: FailableAction) {
194 self.failure.lock().await.on_write = FailureMode::Once(action);
195 }
196
197 pub async fn fail_one_commit(&self, action: FailableAction) {
198 self.failure.lock().await.on_commit = FailureMode::Once(action);
199 }
200
201 pub async fn fail_one_begin_writable(&self, action: FailableAction) {
202 self.failure.lock().await.on_begin_writable = FailureMode::Once(action);
203 }
204
205 pub async fn fail_one_begin_read_only(&self, action: FailableAction) {
206 self.failure.lock().await.on_begin_read_only = FailureMode::Once(action);
207 }
208}
209
210impl<S> VersionedDataSource for FailStorage<S>
211where
212 S: VersionedDataSource,
213{
214 type Transaction<'a>
215 = Transaction<S::Transaction<'a>>
216 where
217 Self: 'a;
218 type ReadOnly<'a>
219 = Transaction<S::ReadOnly<'a>>
220 where
221 Self: 'a;
222
223 async fn write(&self) -> anyhow::Result<<Self as VersionedDataSource>::Transaction<'_>> {
224 self.failure
225 .lock()
226 .await
227 .on_begin_writable
228 .maybe_fail(FailableAction::Any)?;
229 Ok(Transaction {
230 inner: self.inner.write().await?,
231 failure: self.failure.clone(),
232 })
233 }
234
235 async fn read(&self) -> anyhow::Result<<Self as VersionedDataSource>::ReadOnly<'_>> {
236 self.failure
237 .lock()
238 .await
239 .on_begin_read_only
240 .maybe_fail(FailableAction::Any)?;
241 Ok(Transaction {
242 inner: self.inner.read().await?,
243 failure: self.failure.clone(),
244 })
245 }
246}
247
248impl<S> PrunerConfig for FailStorage<S>
249where
250 S: PrunerConfig,
251{
252 fn set_pruning_config(&mut self, cfg: PrunerCfg) {
253 self.inner.set_pruning_config(cfg);
254 }
255
256 fn get_pruning_config(&self) -> Option<PrunerCfg> {
257 self.inner.get_pruning_config()
258 }
259}
260
261#[async_trait]
262impl<S> PruneStorage for FailStorage<S>
263where
264 S: PruneStorage + Sync,
265{
266 type Pruner<'a>
267 = S::Pruner<'a>
268 where
269 S: 'a;
270
271 async fn prune<'a>(&'a self, pruner: &mut Self::Pruner<'a>) -> anyhow::Result<Option<u64>> {
272 self.inner.prune(pruner).await
273 }
274}
275
276impl<S> HasMetrics for FailStorage<S>
277where
278 S: HasMetrics,
279{
280 fn metrics(&self) -> &PrometheusMetrics {
281 self.inner.metrics()
282 }
283}
284
285#[derive(Debug)]
286pub struct Transaction<T> {
287 inner: T,
288 failure: Arc<Mutex<Failure>>,
289}
290
291impl<T> Transaction<T> {
292 async fn maybe_fail_read(&self, action: FailableAction) -> QueryResult<()> {
293 self.failure.lock().await.on_read.maybe_fail(action)
294 }
295
296 async fn maybe_fail_write(&self, action: FailableAction) -> QueryResult<()> {
297 self.failure.lock().await.on_write.maybe_fail(action)
298 }
299
300 async fn maybe_fail_commit(&self, action: FailableAction) -> QueryResult<()> {
301 self.failure.lock().await.on_commit.maybe_fail(action)
302 }
303}
304
305impl<T> update::Transaction for Transaction<T>
306where
307 T: update::Transaction,
308{
309 async fn commit(self) -> anyhow::Result<()> {
310 self.maybe_fail_commit(FailableAction::Any).await?;
311 self.inner.commit().await
312 }
313
314 fn revert(self) -> impl Future + Send {
315 self.inner.revert()
316 }
317}
318
319#[async_trait]
320impl<Types, T> AvailabilityStorage<Types> for Transaction<T>
321where
322 Types: NodeType,
323 Header<Types>: QueryableHeader<Types>,
324 Payload<Types>: QueryablePayload<Types>,
325 T: AvailabilityStorage<Types>,
326{
327 async fn get_leaf(&mut self, id: LeafId<Types>) -> QueryResult<LeafQueryData<Types>> {
328 self.maybe_fail_read(FailableAction::GetLeaf).await?;
329 self.inner.get_leaf(id).await
330 }
331
332 async fn get_block(&mut self, id: BlockId<Types>) -> QueryResult<BlockQueryData<Types>> {
333 self.maybe_fail_read(FailableAction::GetBlock).await?;
334 self.inner.get_block(id).await
335 }
336
337 async fn get_header(&mut self, id: BlockId<Types>) -> QueryResult<Header<Types>> {
338 self.maybe_fail_read(FailableAction::GetHeader).await?;
339 self.inner.get_header(id).await
340 }
341
342 async fn get_payload(&mut self, id: BlockId<Types>) -> QueryResult<PayloadQueryData<Types>> {
343 self.maybe_fail_read(FailableAction::GetPayload).await?;
344 self.inner.get_payload(id).await
345 }
346
347 async fn get_payload_metadata(
348 &mut self,
349 id: BlockId<Types>,
350 ) -> QueryResult<PayloadMetadata<Types>> {
351 self.maybe_fail_read(FailableAction::GetPayloadMetadata)
352 .await?;
353 self.inner.get_payload_metadata(id).await
354 }
355
356 async fn get_vid_common(
357 &mut self,
358 id: BlockId<Types>,
359 ) -> QueryResult<VidCommonQueryData<Types>> {
360 self.maybe_fail_read(FailableAction::GetVidCommon).await?;
361 self.inner.get_vid_common(id).await
362 }
363
364 async fn get_vid_common_metadata(
365 &mut self,
366 id: BlockId<Types>,
367 ) -> QueryResult<VidCommonMetadata<Types>> {
368 self.maybe_fail_read(FailableAction::GetVidCommonMetadata)
369 .await?;
370 self.inner.get_vid_common_metadata(id).await
371 }
372
373 async fn get_leaf_range<R>(
374 &mut self,
375 range: R,
376 ) -> QueryResult<Vec<QueryResult<LeafQueryData<Types>>>>
377 where
378 R: RangeBounds<usize> + Send + 'static,
379 {
380 self.maybe_fail_read(FailableAction::GetLeafRange).await?;
381 self.inner.get_leaf_range(range).await
382 }
383
384 async fn get_block_range<R>(
385 &mut self,
386 range: R,
387 ) -> QueryResult<Vec<QueryResult<BlockQueryData<Types>>>>
388 where
389 R: RangeBounds<usize> + Send + 'static,
390 {
391 self.maybe_fail_read(FailableAction::GetBlockRange).await?;
392 self.inner.get_block_range(range).await
393 }
394
395 async fn get_payload_range<R>(
396 &mut self,
397 range: R,
398 ) -> QueryResult<Vec<QueryResult<PayloadQueryData<Types>>>>
399 where
400 R: RangeBounds<usize> + Send + 'static,
401 {
402 self.maybe_fail_read(FailableAction::GetPayloadRange)
403 .await?;
404 self.inner.get_payload_range(range).await
405 }
406
407 async fn get_payload_metadata_range<R>(
408 &mut self,
409 range: R,
410 ) -> QueryResult<Vec<QueryResult<PayloadMetadata<Types>>>>
411 where
412 R: RangeBounds<usize> + Send + 'static,
413 {
414 self.maybe_fail_read(FailableAction::GetPayloadMetadataRange)
415 .await?;
416 self.inner.get_payload_metadata_range(range).await
417 }
418
419 async fn get_vid_common_range<R>(
420 &mut self,
421 range: R,
422 ) -> QueryResult<Vec<QueryResult<VidCommonQueryData<Types>>>>
423 where
424 R: RangeBounds<usize> + Send + 'static,
425 {
426 self.maybe_fail_read(FailableAction::GetVidCommonRange)
427 .await?;
428 self.inner.get_vid_common_range(range).await
429 }
430
431 async fn get_vid_common_metadata_range<R>(
432 &mut self,
433 range: R,
434 ) -> QueryResult<Vec<QueryResult<VidCommonMetadata<Types>>>>
435 where
436 R: RangeBounds<usize> + Send + 'static,
437 {
438 self.maybe_fail_read(FailableAction::GetVidCommonMetadataRange)
439 .await?;
440 self.inner.get_vid_common_metadata_range(range).await
441 }
442
443 async fn get_block_with_transaction(
444 &mut self,
445 hash: TransactionHash<Types>,
446 ) -> QueryResult<BlockQueryData<Types>> {
447 self.maybe_fail_read(FailableAction::GetTransaction).await?;
448 self.inner.get_block_with_transaction(hash).await
449 }
450}
451
452impl<Types, T> UpdateAvailabilityStorage<Types> for Transaction<T>
453where
454 Types: NodeType,
455 Header<Types>: QueryableHeader<Types>,
456 Payload<Types>: QueryablePayload<Types>,
457 T: UpdateAvailabilityStorage<Types> + Send + Sync,
458{
459 async fn insert_qc_chain(
460 &mut self,
461 height: u64,
462 qc_chain: Option<[CertificatePair<Types>; 2]>,
463 ) -> anyhow::Result<()> {
464 self.maybe_fail_write(FailableAction::Any).await?;
465 self.inner.insert_qc_chain(height, qc_chain).await
466 }
467
468 async fn insert_cert2(
469 &mut self,
470 height: u64,
471 cert2: Certificate2<Types>,
472 ) -> anyhow::Result<()> {
473 self.maybe_fail_write(FailableAction::Any).await?;
474 self.inner.insert_cert2(height, cert2).await
475 }
476
477 async fn insert_leaf_range<'a>(
478 &mut self,
479 leaves: impl Send + IntoIterator<IntoIter: Send, Item = &'a LeafQueryData<Types>>,
480 ) -> anyhow::Result<()> {
481 self.maybe_fail_write(FailableAction::Any).await?;
482 self.inner.insert_leaf_range(leaves).await
483 }
484
485 async fn insert_block_range<'a>(
486 &mut self,
487 blocks: impl Send + IntoIterator<IntoIter: Send, Item = &'a BlockQueryData<Types>>,
488 ) -> anyhow::Result<()> {
489 self.maybe_fail_write(FailableAction::Any).await?;
490 self.inner.insert_block_range(blocks).await
491 }
492
493 async fn insert_vid_range<'a>(
494 &mut self,
495 vid: impl Send
496 + IntoIterator<
497 IntoIter: Send,
498 Item = (&'a VidCommonQueryData<Types>, Option<&'a VidShare>),
499 >,
500 ) -> anyhow::Result<()> {
501 self.maybe_fail_write(FailableAction::Any).await?;
502 self.inner.insert_vid_range(vid).await
503 }
504}
505
506#[async_trait]
507impl<T> PrunedHeightStorage for Transaction<T>
508where
509 T: PrunedHeightStorage + Send + Sync,
510{
511 async fn load_pruned_height(&mut self) -> anyhow::Result<Option<u64>> {
512 self.maybe_fail_read(FailableAction::Any).await?;
513 self.inner.load_pruned_height().await
514 }
515
516 async fn load_state_pruned_height(&mut self) -> anyhow::Result<Option<u64>> {
517 self.maybe_fail_read(FailableAction::Any).await?;
518 self.inner.load_state_pruned_height().await
519 }
520}
521
522#[async_trait]
523impl<Types, T> NodeStorage<Types> for Transaction<T>
524where
525 Types: NodeType,
526 Header<Types>: QueryableHeader<Types>,
527 T: NodeStorage<Types> + Send + Sync,
528{
529 async fn block_height(&mut self) -> QueryResult<usize> {
530 self.maybe_fail_read(FailableAction::Any).await?;
531 self.inner.block_height().await
532 }
533
534 async fn count_transactions_in_range(
535 &mut self,
536 range: impl RangeBounds<usize> + Send,
537 namespace: Option<NamespaceId<Types>>,
538 ) -> QueryResult<usize> {
539 self.maybe_fail_read(FailableAction::Any).await?;
540 self.inner
541 .count_transactions_in_range(range, namespace)
542 .await
543 }
544
545 async fn payload_size_in_range(
546 &mut self,
547 range: impl RangeBounds<usize> + Send,
548 namespace: Option<NamespaceId<Types>>,
549 ) -> QueryResult<usize> {
550 self.maybe_fail_read(FailableAction::Any).await?;
551 self.inner.payload_size_in_range(range, namespace).await
552 }
553
554 async fn vid_share<ID>(&mut self, id: ID) -> QueryResult<VidShare>
555 where
556 ID: Into<BlockId<Types>> + Send + Sync,
557 {
558 self.maybe_fail_read(FailableAction::Any).await?;
559 self.inner.vid_share(id).await
560 }
561
562 async fn sync_status_for_range(
563 &mut self,
564 start: usize,
565 end: usize,
566 ) -> QueryResult<SyncStatusQueryData> {
567 self.maybe_fail_read(FailableAction::Any).await?;
568 self.inner.sync_status_for_range(start, end).await
569 }
570
571 async fn get_header_window(
572 &mut self,
573 start: impl Into<WindowStart<Types>> + Send + Sync,
574 end: u64,
575 limit: usize,
576 ) -> QueryResult<TimeWindowQueryData<Header<Types>>> {
577 self.maybe_fail_read(FailableAction::Any).await?;
578 self.inner.get_header_window(start, end, limit).await
579 }
580
581 async fn latest_qc_chain(&mut self) -> QueryResult<Option<[CertificatePair<Types>; 2]>> {
582 self.maybe_fail_read(FailableAction::Any).await?;
583 self.inner.latest_qc_chain().await
584 }
585
586 async fn load_cert2(&mut self, height: u64) -> QueryResult<Option<Certificate2<Types>>> {
587 self.maybe_fail_read(FailableAction::Any).await?;
588 self.inner.load_cert2(height).await
589 }
590
591 async fn load_earliest_cert2(
592 &mut self,
593 height: u64,
594 ) -> QueryResult<Option<Certificate2<Types>>> {
595 self.maybe_fail_read(FailableAction::Any).await?;
596 self.inner.load_earliest_cert2(height).await
597 }
598}
599
600impl<Types, T> AggregatesStorage<Types> for Transaction<T>
601where
602 Types: NodeType,
603 Header<Types>: QueryableHeader<Types>,
604 T: AggregatesStorage<Types> + Send + Sync,
605{
606 async fn aggregates_height(&mut self) -> anyhow::Result<usize> {
607 self.maybe_fail_read(FailableAction::Any).await?;
608 self.inner.aggregates_height().await
609 }
610
611 async fn load_prev_aggregate(&mut self) -> anyhow::Result<Option<Aggregate<Types>>> {
612 self.maybe_fail_read(FailableAction::Any).await?;
613 self.inner.load_prev_aggregate().await
614 }
615}
616
617impl<T, Types> UpdateAggregatesStorage<Types> for Transaction<T>
618where
619 Types: NodeType,
620 Header<Types>: QueryableHeader<Types>,
621 T: UpdateAggregatesStorage<Types> + Send + Sync,
622{
623 async fn update_aggregates(
624 &mut self,
625 prev: Aggregate<Types>,
626 blocks: &[PayloadMetadata<Types>],
627 ) -> anyhow::Result<Aggregate<Types>> {
628 self.maybe_fail_write(FailableAction::Any).await?;
629 self.inner.update_aggregates(prev, blocks).await
630 }
631}