openzeppelin_relayer/repositories/transaction_counter/
transaction_counter_in_memory.rs

1//! This module provides an in-memory implementation of a transaction counter.
2//!
3//! The `InMemoryTransactionCounter` struct is used to track and manage transaction nonces
4//! for different relayers and addresses. It supports operations to get, increment, decrement,
5//! and set nonce values. This implementation uses a `DashMap` for concurrent access and
6//! modification of the nonce values.
7use 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>, // (relayer_id, address) -> nonce/sequence
15}
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        // Hold the shard entry lock across read-compare-write so the monotonic max is atomic
77        // with respect to concurrent `sync_floor`/`set` on the same key.
78        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        // Initially should be None
112        assert_eq!(store.get(relayer_id, address).await.unwrap(), None);
113
114        // Set a value explicitly
115        store.set(relayer_id, address, 100).await.unwrap();
116        assert_eq!(store.get(relayer_id, address).await.unwrap(), Some(100));
117
118        // Increment
119        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        // Decrement
126        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        // Counter starts at sequence S = 5 (next to allocate).
136        store.set(relayer, address, 5).await.unwrap();
137        // Concurrent allocations advance it 10 ahead -> 15.
138        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        // Recovery observes the stale chain floor (S+1 = 6); it MUST NOT rewind below 15.
144        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        // But it DOES advance when the chain floor is genuinely ahead.
149        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        // Setup different relayer/address combinations
184        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        // Verify independent counters
189        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        // Verify independent increments
194        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}