openzeppelin_relayer/repositories/transaction_counter/
mod.rs1pub 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 async fn sync_floor(
70 &self,
71 relayer_id: &str,
72 address: &str,
73 floor: u64,
74 ) -> Result<u64, RepositoryError>;
75
76 async fn drop_all_entries(&self) -> Result<(), RepositoryError>;
79}
80
81#[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 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}