openzeppelin_relayer/repositories/transaction_counter/
transaction_counter_in_memory.rs1use async_trait::async_trait;
8use dashmap::DashMap;
9
10use crate::repositories::{RepositoryError, TransactionCounterTrait};
11
12#[derive(Debug, Default, Clone)]
13pub struct InMemoryTransactionCounter {
14 store: DashMap<(String, String), u64>, }
16
17impl InMemoryTransactionCounter {
18 pub fn new() -> Self {
19 Self {
20 store: DashMap::new(),
21 }
22 }
23}
24
25#[async_trait]
26impl TransactionCounterTrait for InMemoryTransactionCounter {
27 async fn get(&self, relayer_id: &str, address: &str) -> Result<Option<u64>, RepositoryError> {
28 Ok(self
29 .store
30 .get(&(relayer_id.to_string(), address.to_string()))
31 .map(|n| *n))
32 }
33
34 async fn get_and_increment(
35 &self,
36 relayer_id: &str,
37 address: &str,
38 ) -> Result<u64, RepositoryError> {
39 let mut entry = self
40 .store
41 .entry((relayer_id.to_string(), address.to_string()))
42 .or_insert(0);
43 let current = *entry;
44 *entry += 1;
45 Ok(current)
46 }
47
48 async fn decrement(&self, relayer_id: &str, address: &str) -> Result<u64, RepositoryError> {
49 let mut entry = self
50 .store
51 .get_mut(&(relayer_id.to_string(), address.to_string()))
52 .ok_or_else(|| RepositoryError::NotFound(format!("Counter not found for {address}")))?;
53 if *entry > 0 {
54 *entry -= 1;
55 }
56 Ok(*entry)
57 }
58
59 async fn set(
60 &self,
61 relayer_id: &str,
62 address: &str,
63 value: u64,
64 ) -> Result<(), RepositoryError> {
65 self.store
66 .insert((relayer_id.to_string(), address.to_string()), value);
67 Ok(())
68 }
69
70 async fn sync_floor(
71 &self,
72 relayer_id: &str,
73 address: &str,
74 floor: u64,
75 ) -> Result<u64, RepositoryError> {
76 let mut entry = self
79 .store
80 .entry((relayer_id.to_string(), address.to_string()))
81 .or_insert(floor);
82 if *entry < floor {
83 *entry = floor;
84 }
85 Ok(*entry)
86 }
87
88 async fn drop_all_entries(&self) -> Result<(), RepositoryError> {
89 self.store.clear();
90 Ok(())
91 }
92}
93
94#[cfg(test)]
95mod tests {
96 use super::*;
97
98 #[tokio::test]
99 async fn test_decrement_not_found() {
100 let store = InMemoryTransactionCounter::new();
101 let result = store.decrement("nonexistent", "0x1234").await;
102 assert!(matches!(result, Err(RepositoryError::NotFound(_))));
103 }
104
105 #[tokio::test]
106 async fn test_nonce_store() {
107 let store = InMemoryTransactionCounter::new();
108 let relayer_id = "relayer_1";
109 let address = "0x1234";
110
111 assert_eq!(store.get(relayer_id, address).await.unwrap(), None);
113
114 store.set(relayer_id, address, 100).await.unwrap();
116 assert_eq!(store.get(relayer_id, address).await.unwrap(), Some(100));
117
118 assert_eq!(
120 store.get_and_increment(relayer_id, address).await.unwrap(),
121 100
122 );
123 assert_eq!(store.get(relayer_id, address).await.unwrap(), Some(101));
124
125 assert_eq!(store.decrement(relayer_id, address).await.unwrap(), 100);
127 assert_eq!(store.get(relayer_id, address).await.unwrap(), Some(100));
128 }
129
130 #[tokio::test]
131 async fn test_sync_floor_no_rewind() {
132 let store = InMemoryTransactionCounter::new();
133 let (relayer, address) = ("relayer_1", "0xabc");
134
135 store.set(relayer, address, 5).await.unwrap();
137 for _ in 0..10 {
139 store.get_and_increment(relayer, address).await.unwrap();
140 }
141 assert_eq!(store.get(relayer, address).await.unwrap(), Some(15));
142
143 let effective = store.sync_floor(relayer, address, 6).await.unwrap();
145 assert_eq!(effective, 15);
146 assert_eq!(store.get(relayer, address).await.unwrap(), Some(15));
147
148 let effective = store.sync_floor(relayer, address, 20).await.unwrap();
150 assert_eq!(effective, 20);
151 assert_eq!(store.get(relayer, address).await.unwrap(), Some(20));
152 }
153
154 #[tokio::test]
155 async fn test_sync_floor_seeds_when_unset() {
156 let store = InMemoryTransactionCounter::new();
157 let effective = store.sync_floor("relayer_1", "0xabc", 7).await.unwrap();
158 assert_eq!(effective, 7);
159 assert_eq!(store.get("relayer_1", "0xabc").await.unwrap(), Some(7));
160 }
161
162 #[tokio::test]
163 async fn test_drop_all_entries() {
164 let store = InMemoryTransactionCounter::new();
165
166 store.set("relayer_1", "0x1234", 100).await.unwrap();
167 store.set("relayer_1", "0x5678", 200).await.unwrap();
168 store.set("relayer_2", "0x1234", 300).await.unwrap();
169
170 assert_eq!(store.get("relayer_1", "0x1234").await.unwrap(), Some(100));
171
172 store.drop_all_entries().await.unwrap();
173
174 assert_eq!(store.get("relayer_1", "0x1234").await.unwrap(), None);
175 assert_eq!(store.get("relayer_1", "0x5678").await.unwrap(), None);
176 assert_eq!(store.get("relayer_2", "0x1234").await.unwrap(), None);
177 }
178
179 #[tokio::test]
180 async fn test_multiple_relayers() {
181 let store = InMemoryTransactionCounter::new();
182
183 store.set("relayer_1", "0x1234", 100).await.unwrap();
185 store.set("relayer_1", "0x5678", 200).await.unwrap();
186 store.set("relayer_2", "0x1234", 300).await.unwrap();
187
188 assert_eq!(store.get("relayer_1", "0x1234").await.unwrap(), Some(100));
190 assert_eq!(store.get("relayer_1", "0x5678").await.unwrap(), Some(200));
191 assert_eq!(store.get("relayer_2", "0x1234").await.unwrap(), Some(300));
192
193 assert_eq!(
195 store
196 .get_and_increment("relayer_1", "0x1234")
197 .await
198 .unwrap(),
199 100
200 );
201 assert_eq!(
202 store
203 .get_and_increment("relayer_1", "0x1234")
204 .await
205 .unwrap(),
206 101
207 );
208 assert_eq!(
209 store
210 .get_and_increment("relayer_1", "0x5678")
211 .await
212 .unwrap(),
213 200
214 );
215 assert_eq!(
216 store
217 .get_and_increment("relayer_1", "0x5678")
218 .await
219 .unwrap(),
220 201
221 );
222 assert_eq!(store.get("relayer_2", "0x1234").await.unwrap(), Some(300));
223 }
224}