Skip to main content

hotshot_query_service/data_source/storage/
fail_storage.rs

1// Copyright (c) 2022 Espresso Systems (espressosys.com)
2// This file is part of the HotShot Query Service library.
3//
4// This program is free software: you can redistribute it and/or modify it under the terms of the GNU
5// General Public License as published by the Free Software Foundation, either version 3 of the
6// License, or (at your option) any later version.
7// This program is distributed in the hope that it will be useful, but WITHOUT ANY WARRANTY; without
8// even the implied warranty of MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the GNU
9// General Public License for more details.
10// You should have received a copy of the GNU General Public License along with this program. If not,
11// see <https://www.gnu.org/licenses/>.
12
13#![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/// A specific action that can be targeted to inject an error.
46#[derive(Clone, Copy, Debug, PartialEq, Eq)]
47pub enum FailableAction {
48    // TODO currently we implement failable actions for the availability methods, but if needed we
49    // can always add more variants for other actions.
50    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    /// Target any action for failure.
69    Any,
70}
71
72impl FailableAction {
73    /// Should `self` being targeted for failure cause `action` to fail?
74    fn matches(self, action: Self) -> bool {
75        // Fail if this is the action specifically targeted for failure or if we are failing any
76        // action right now.
77        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/// Storage wrapper for error injection.
115#[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}