openzeppelin_relayer/repositories/transaction_counter/
mod.rs

1//! Transaction Counter Repository Module
2//!
3//! This module provides the transaction counter repository layer for the OpenZeppelin Relayer service.
4//! It implements specialized counters for tracking transaction nonces and sequence numbers
5//! across different blockchain networks, supporting both in-memory and Redis-backed storage.
6//!
7//! ## Repository Implementations
8//!
9//! - [`InMemoryTransactionCounter`]: Fast in-memory storage using DashMap for concurrency
10//! - [`RedisTransactionCounter`]: Redis-backed storage for production environments
11//!
12//! ## Counter Operations
13//!
14//! The transaction counter supports several key operations:
15//!
16//! - **Get**: Retrieve current counter value
17//! - **Get and Increment**: Atomically get current value and increment
18//! - **Decrement**: Decrement counter (for rollbacks)
19//! - **Set**: Set counter to specific value
20//!
21pub mod transaction_counter_in_memory;
22pub mod transaction_counter_redis;
23
24use crate::utils::RedisConnections;
25pub use transaction_counter_in_memory::InMemoryTransactionCounter;
26pub use transaction_counter_redis::RedisTransactionCounter;
27
28use async_trait::async_trait;
29use serde::Serialize;
30use std::sync::Arc;
31use thiserror::Error;
32
33#[cfg(test)]
34use mockall::automock;
35
36use crate::models::RepositoryError;
37
38#[derive(Error, Debug, Serialize)]
39pub enum TransactionCounterError {
40    #[error("No sequence found for relayer {relayer_id} and address {address}")]
41    SequenceNotFound { relayer_id: String, address: String },
42    #[error("Counter not found for {0}")]
43    NotFound(String),
44}
45
46#[allow(dead_code)]
47#[async_trait]
48#[cfg_attr(test, automock)]
49pub trait TransactionCounterTrait {
50    async fn get(&self, relayer_id: &str, address: &str) -> Result<Option<u64>, RepositoryError>;
51
52    async fn get_and_increment(
53        &self,
54        relayer_id: &str,
55        address: &str,
56    ) -> Result<u64, RepositoryError>;
57
58    async fn decrement(&self, relayer_id: &str, address: &str) -> Result<u64, RepositoryError>;
59
60    async fn set(&self, relayer_id: &str, address: &str, value: u64)
61        -> Result<(), RepositoryError>;
62
63    /// Monotonically raise the counter to `floor` only if the current value is lower.
64    ///
65    /// Used for race-safe bad-sequence recovery: the live chain sequence is treated as a
66    /// floor, never as the authoritative next assignment, so concurrent allocations that
67    /// have already advanced the counter beyond `floor` are never rewound. Returns the
68    /// effective value after the operation (always `>= floor`).
69    async fn sync_floor(
70        &self,
71        relayer_id: &str,
72        address: &str,
73        floor: u64,
74    ) -> Result<u64, RepositoryError>;
75
76    /// Remove all stored counter entries from the underlying backend.
77    /// Intended for startup reset flows when `RESET_STORAGE_ON_START` is enabled.
78    async fn drop_all_entries(&self) -> Result<(), RepositoryError>;
79}
80
81/// Enum wrapper for different transaction counter repository implementations
82#[derive(Debug, Clone)]
83pub enum TransactionCounterRepositoryStorage {
84    InMemory(InMemoryTransactionCounter),
85    Redis(RedisTransactionCounter),
86}
87
88impl TransactionCounterRepositoryStorage {
89    pub fn new_in_memory() -> Self {
90        Self::InMemory(InMemoryTransactionCounter::new())
91    }
92    pub fn new_redis(
93        connections: Arc<RedisConnections>,
94        key_prefix: String,
95    ) -> Result<Self, RepositoryError> {
96        Ok(Self::Redis(RedisTransactionCounter::new(
97            connections,
98            key_prefix,
99        )?))
100    }
101}
102
103#[async_trait]
104impl TransactionCounterTrait for TransactionCounterRepositoryStorage {
105    async fn get(&self, relayer_id: &str, address: &str) -> Result<Option<u64>, RepositoryError> {
106        match self {
107            TransactionCounterRepositoryStorage::InMemory(counter) => {
108                counter.get(relayer_id, address).await
109            }
110            TransactionCounterRepositoryStorage::Redis(counter) => {
111                counter.get(relayer_id, address).await
112            }
113        }
114    }
115
116    async fn get_and_increment(
117        &self,
118        relayer_id: &str,
119        address: &str,
120    ) -> Result<u64, RepositoryError> {
121        match self {
122            TransactionCounterRepositoryStorage::InMemory(counter) => {
123                counter.get_and_increment(relayer_id, address).await
124            }
125            TransactionCounterRepositoryStorage::Redis(counter) => {
126                counter.get_and_increment(relayer_id, address).await
127            }
128        }
129    }
130
131    async fn decrement(&self, relayer_id: &str, address: &str) -> Result<u64, RepositoryError> {
132        match self {
133            TransactionCounterRepositoryStorage::InMemory(counter) => {
134                counter.decrement(relayer_id, address).await
135            }
136            TransactionCounterRepositoryStorage::Redis(counter) => {
137                counter.decrement(relayer_id, address).await
138            }
139        }
140    }
141
142    async fn set(
143        &self,
144        relayer_id: &str,
145        address: &str,
146        value: u64,
147    ) -> Result<(), RepositoryError> {
148        match self {
149            TransactionCounterRepositoryStorage::InMemory(counter) => {
150                counter.set(relayer_id, address, value).await
151            }
152            TransactionCounterRepositoryStorage::Redis(counter) => {
153                counter.set(relayer_id, address, value).await
154            }
155        }
156    }
157
158    async fn sync_floor(
159        &self,
160        relayer_id: &str,
161        address: &str,
162        floor: u64,
163    ) -> Result<u64, RepositoryError> {
164        match self {
165            TransactionCounterRepositoryStorage::InMemory(counter) => {
166                counter.sync_floor(relayer_id, address, floor).await
167            }
168            TransactionCounterRepositoryStorage::Redis(counter) => {
169                counter.sync_floor(relayer_id, address, floor).await
170            }
171        }
172    }
173
174    async fn drop_all_entries(&self) -> Result<(), RepositoryError> {
175        match self {
176            TransactionCounterRepositoryStorage::InMemory(counter) => {
177                counter.drop_all_entries().await
178            }
179            TransactionCounterRepositoryStorage::Redis(counter) => counter.drop_all_entries().await,
180        }
181    }
182}
183
184#[cfg(test)]
185mod tests {
186
187    use super::*;
188
189    #[tokio::test]
190    async fn test_in_memory_repository_creation() {
191        let repo = TransactionCounterRepositoryStorage::new_in_memory();
192
193        matches!(repo, TransactionCounterRepositoryStorage::InMemory(_));
194    }
195
196    #[tokio::test]
197    async fn test_enum_wrapper_delegation() {
198        let repo = TransactionCounterRepositoryStorage::new_in_memory();
199
200        // Test that the enum wrapper properly delegates to the underlying implementation
201        let result = repo.get("test_relayer", "0x1234").await.unwrap();
202        assert_eq!(result, None);
203
204        repo.set("test_relayer", "0x1234", 100).await.unwrap();
205        let result = repo.get("test_relayer", "0x1234").await.unwrap();
206        assert_eq!(result, Some(100));
207
208        let current = repo
209            .get_and_increment("test_relayer", "0x1234")
210            .await
211            .unwrap();
212        assert_eq!(current, 100);
213
214        let result = repo.get("test_relayer", "0x1234").await.unwrap();
215        assert_eq!(result, Some(101));
216
217        let new_value = repo.decrement("test_relayer", "0x1234").await.unwrap();
218        assert_eq!(new_value, 100);
219    }
220
221    #[tokio::test]
222    async fn test_enum_wrapper_drop_all_entries() {
223        let repo = TransactionCounterRepositoryStorage::new_in_memory();
224
225        repo.set("relayer_1", "0x1234", 100).await.unwrap();
226        repo.set("relayer_2", "0x5678", 200).await.unwrap();
227
228        repo.drop_all_entries().await.unwrap();
229
230        assert_eq!(repo.get("relayer_1", "0x1234").await.unwrap(), None);
231        assert_eq!(repo.get("relayer_2", "0x5678").await.unwrap(), None);
232    }
233}