openzeppelin_relayer/repositories/transaction_counter/
transaction_counter_redis.rs

1//! Redis implementation of the transaction counter.
2//!
3//! This module provides a Redis-based implementation of the `TransactionCounterTrait`,
4//! allowing transaction counters to be stored and retrieved from a Redis database.
5//! The implementation includes comprehensive error handling, logging, and atomic operations
6//! to ensure consistency when incrementing and decrementing counters.
7
8use super::TransactionCounterTrait;
9use crate::models::RepositoryError;
10use crate::repositories::redis_base::RedisRepository;
11use crate::utils::RedisConnections;
12use async_trait::async_trait;
13use redis::AsyncCommands;
14use std::fmt;
15use std::sync::Arc;
16use tracing::debug;
17
18const COUNTER_PREFIX: &str = "transaction_counter";
19
20#[derive(Clone)]
21pub struct RedisTransactionCounter {
22    pub connections: Arc<RedisConnections>,
23    pub key_prefix: String,
24}
25
26impl RedisRepository for RedisTransactionCounter {}
27
28impl fmt::Debug for RedisTransactionCounter {
29    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
30        f.debug_struct("RedisTransactionCounter")
31            .field("key_prefix", &self.key_prefix)
32            .finish()
33    }
34}
35
36impl RedisTransactionCounter {
37    pub fn new(
38        connections: Arc<RedisConnections>,
39        key_prefix: String,
40    ) -> Result<Self, RepositoryError> {
41        if key_prefix.is_empty() {
42            return Err(RepositoryError::InvalidData(
43                "Redis key prefix cannot be empty".to_string(),
44            ));
45        }
46
47        Ok(Self {
48            connections,
49            key_prefix,
50        })
51    }
52
53    /// Generate key for transaction counter: {prefix}:transaction_counter:{relayer_id}:{address}
54    fn counter_key(&self, relayer_id: &str, address: &str) -> String {
55        format!(
56            "{}:{}:{}:{}",
57            self.key_prefix, COUNTER_PREFIX, relayer_id, address
58        )
59    }
60}
61
62#[async_trait]
63impl TransactionCounterTrait for RedisTransactionCounter {
64    async fn get(&self, relayer_id: &str, address: &str) -> Result<Option<u64>, RepositoryError> {
65        if relayer_id.is_empty() {
66            return Err(RepositoryError::InvalidData(
67                "Relayer ID cannot be empty".to_string(),
68            ));
69        }
70
71        if address.is_empty() {
72            return Err(RepositoryError::InvalidData(
73                "Address cannot be empty".to_string(),
74            ));
75        }
76
77        let key = self.counter_key(relayer_id, address);
78        debug!(relayer_id = %relayer_id, address = %address, "getting counter for relayer and address");
79
80        let mut conn = self
81            .get_connection(self.connections.reader(), "get")
82            .await?;
83
84        let value: Option<u64> = conn
85            .get(&key)
86            .await
87            .map_err(|e| self.map_redis_error(e, "get_counter"))?;
88
89        debug!(value = ?value, "retrieved counter value");
90        Ok(value)
91    }
92
93    async fn get_and_increment(
94        &self,
95        relayer_id: &str,
96        address: &str,
97    ) -> Result<u64, RepositoryError> {
98        if relayer_id.is_empty() {
99            return Err(RepositoryError::InvalidData(
100                "Relayer ID cannot be empty".to_string(),
101            ));
102        }
103
104        if address.is_empty() {
105            return Err(RepositoryError::InvalidData(
106                "Address cannot be empty".to_string(),
107            ));
108        }
109
110        let key = self.counter_key(relayer_id, address);
111        debug!(relayer_id = %relayer_id, address = %address, "getting and incrementing counter for relayer and address");
112
113        let mut conn = self
114            .get_connection(self.connections.primary(), "get_and_increment")
115            .await?;
116
117        // Use Redis INCR for atomic increment
118        let new_value: u64 = conn
119            .incr(&key, 1)
120            .await
121            .map_err(|e| self.map_redis_error(e, "get_and_increment"))?;
122
123        let current = new_value.saturating_sub(1);
124
125        debug!(from = %current, to = %(current + 1), "counter incremented");
126        Ok(current)
127    }
128
129    async fn decrement(&self, relayer_id: &str, address: &str) -> Result<u64, RepositoryError> {
130        if relayer_id.is_empty() {
131            return Err(RepositoryError::InvalidData(
132                "Relayer ID cannot be empty".to_string(),
133            ));
134        }
135
136        if address.is_empty() {
137            return Err(RepositoryError::InvalidData(
138                "Address cannot be empty".to_string(),
139            ));
140        }
141
142        let key = self.counter_key(relayer_id, address);
143        debug!(relayer_id = %relayer_id, address = %address, "decrementing counter for relayer and address");
144
145        let mut conn = self
146            .get_connection(self.connections.primary(), "decrement")
147            .await?;
148
149        // Check if counter exists first
150        let exists: bool = conn
151            .exists(&key)
152            .await
153            .map_err(|e| self.map_redis_error(e, "check_counter_exists"))?;
154
155        if !exists {
156            return Err(RepositoryError::NotFound(format!(
157                "Counter not found for relayer {relayer_id} and address {address}"
158            )));
159        }
160
161        // Use Redis DECR and correct if it goes below 0
162        let new_value: i64 = conn
163            .decr(&key, 1)
164            .await
165            .map_err(|e| self.map_redis_error(e, "decrement_counter"))?;
166
167        let new_value = if new_value < 0 {
168            // Correct negative values back to 0
169            let _: () = conn
170                .set(&key, 0)
171                .await
172                .map_err(|e| self.map_redis_error(e, "correct_negative_counter"))?;
173            0u64
174        } else {
175            new_value as u64
176        };
177
178        debug!(new_value = %new_value, "counter decremented");
179        Ok(new_value)
180    }
181
182    async fn set(
183        &self,
184        relayer_id: &str,
185        address: &str,
186        value: u64,
187    ) -> Result<(), RepositoryError> {
188        if relayer_id.is_empty() {
189            return Err(RepositoryError::InvalidData(
190                "Relayer ID cannot be empty".to_string(),
191            ));
192        }
193
194        if address.is_empty() {
195            return Err(RepositoryError::InvalidData(
196                "Address cannot be empty".to_string(),
197            ));
198        }
199
200        let key = self.counter_key(relayer_id, address);
201        debug!(relayer_id = %relayer_id, address = %address, value = %value, "setting counter for relayer and address");
202
203        let mut conn = self
204            .get_connection(self.connections.primary(), "set")
205            .await?;
206
207        let _: () = conn
208            .set(&key, value)
209            .await
210            .map_err(|e| self.map_redis_error(e, "set_counter"))?;
211
212        debug!(value = %value, "counter set");
213        Ok(())
214    }
215
216    async fn sync_floor(
217        &self,
218        relayer_id: &str,
219        address: &str,
220        floor: u64,
221    ) -> Result<u64, RepositoryError> {
222        if relayer_id.is_empty() {
223            return Err(RepositoryError::InvalidData(
224                "Relayer ID cannot be empty".to_string(),
225            ));
226        }
227
228        if address.is_empty() {
229            return Err(RepositoryError::InvalidData(
230                "Address cannot be empty".to_string(),
231            ));
232        }
233
234        let key = self.counter_key(relayer_id, address);
235        debug!(relayer_id = %relayer_id, address = %address, floor = %floor, "syncing counter floor for relayer and address");
236
237        let mut conn = self
238            .get_connection(self.connections.primary(), "sync_floor")
239            .await?;
240
241        // Atomic monotonic max: raise to `floor` only if the current value is lower (or unset),
242        // never rewind below an already-allocated sequence. Returns the effective value.
243        const SYNC_FLOOR_LUA: &str = r#"
244            local cur = redis.call('GET', KEYS[1])
245            if (not cur) or (tonumber(cur) < tonumber(ARGV[1])) then
246                redis.call('SET', KEYS[1], ARGV[1])
247                return tonumber(ARGV[1])
248            end
249            return tonumber(cur)
250        "#;
251
252        let effective: u64 = redis::Script::new(SYNC_FLOOR_LUA)
253            .key(&key)
254            .arg(floor)
255            .invoke_async(&mut conn)
256            .await
257            .map_err(|e| self.map_redis_error(e, "sync_floor"))?;
258
259        debug!(effective = %effective, "counter floor synced");
260        Ok(effective)
261    }
262
263    async fn drop_all_entries(&self) -> Result<(), RepositoryError> {
264        let mut conn = self
265            .get_connection(self.connections.primary(), "drop_all_entries")
266            .await?;
267
268        let pattern = format!("{}:{}:*", self.key_prefix, COUNTER_PREFIX);
269        debug!(pattern = %pattern, "dropping all transaction counter entries");
270
271        // Phase 1: Collect all matching keys without mutating the keyspace.
272        // Deleting during SCAN can cause hash table rehashing, which may skip keys.
273        let mut cursor: u64 = 0;
274        let mut all_keys: Vec<String> = Vec::new();
275
276        loop {
277            let (next_cursor, keys): (u64, Vec<String>) = redis::cmd("SCAN")
278                .cursor_arg(cursor)
279                .arg("MATCH")
280                .arg(&pattern)
281                .arg("COUNT")
282                .arg(100)
283                .query_async(&mut conn)
284                .await
285                .map_err(|e| self.map_redis_error(e, "drop_all_entries_scan"))?;
286
287            all_keys.extend(keys);
288
289            cursor = next_cursor;
290            if cursor == 0 {
291                break;
292            }
293        }
294
295        // Phase 2: Batch delete all collected keys.
296        if !all_keys.is_empty() {
297            let mut pipe = redis::pipe();
298            pipe.atomic();
299            for key in &all_keys {
300                pipe.del(key);
301            }
302            pipe.exec_async(&mut conn)
303                .await
304                .map_err(|e| self.map_redis_error(e, "drop_all_entries_delete"))?;
305        }
306
307        debug!(total_deleted = %all_keys.len(), "dropped all transaction counter entries");
308        Ok(())
309    }
310}
311
312#[cfg(test)]
313mod tests {
314    use super::*;
315    use std::sync::Arc;
316    use tokio;
317    use uuid::Uuid;
318
319    async fn setup_test_repo() -> RedisTransactionCounter {
320        setup_test_repo_with_prefix("test_counter").await
321    }
322
323    async fn setup_test_repo_with_prefix(prefix: &str) -> RedisTransactionCounter {
324        let redis_url =
325            std::env::var("REDIS_URL").unwrap_or_else(|_| "redis://127.0.0.1:6379".to_string());
326        let cfg = deadpool_redis::Config::from_url(&redis_url);
327        let pool = Arc::new(
328            cfg.builder()
329                .expect("Failed to create pool builder")
330                .max_size(16)
331                .runtime(deadpool_redis::Runtime::Tokio1)
332                .build()
333                .expect("Failed to build Redis pool"),
334        );
335        let connections = Arc::new(RedisConnections::new_single_pool(pool));
336
337        RedisTransactionCounter::new(connections, prefix.to_string())
338            .expect("Failed to create Redis transaction counter")
339    }
340
341    #[tokio::test]
342    #[ignore = "Requires active Redis instance"]
343    async fn test_get_nonexistent_counter() {
344        let repo = setup_test_repo().await;
345        let random_id = Uuid::new_v4().to_string();
346        let result = repo.get(&random_id, "0x1234").await.unwrap();
347        assert_eq!(result, None);
348    }
349
350    #[tokio::test]
351    #[ignore = "Requires active Redis instance"]
352    async fn test_set_and_get_counter() {
353        let repo = setup_test_repo().await;
354        let relayer_id = uuid::Uuid::new_v4().to_string();
355        let address = uuid::Uuid::new_v4().to_string();
356
357        repo.set(&relayer_id, &address, 100).await.unwrap();
358        let result = repo.get(&relayer_id, &address).await.unwrap();
359        assert_eq!(result, Some(100));
360    }
361
362    #[tokio::test]
363    #[ignore = "Requires active Redis instance"]
364    async fn test_sync_floor_no_rewind() {
365        let repo = setup_test_repo().await;
366        let relayer_id = uuid::Uuid::new_v4().to_string();
367        let address = uuid::Uuid::new_v4().to_string();
368
369        // Counter starts at S = 5 then concurrent INCRs advance it to S+10 = 15.
370        repo.set(&relayer_id, &address, 5).await.unwrap();
371        for _ in 0..10 {
372            repo.get_and_increment(&relayer_id, &address).await.unwrap();
373        }
374        assert_eq!(repo.get(&relayer_id, &address).await.unwrap(), Some(15));
375
376        // Stale chain floor S+1 = 6 MUST NOT rewind the counter.
377        let effective = repo.sync_floor(&relayer_id, &address, 6).await.unwrap();
378        assert_eq!(effective, 15);
379        assert_eq!(repo.get(&relayer_id, &address).await.unwrap(), Some(15));
380
381        // A genuinely-ahead floor DOES advance the counter.
382        let effective = repo.sync_floor(&relayer_id, &address, 20).await.unwrap();
383        assert_eq!(effective, 20);
384        assert_eq!(repo.get(&relayer_id, &address).await.unwrap(), Some(20));
385    }
386
387    #[tokio::test]
388    #[ignore = "Requires active Redis instance"]
389    async fn test_sync_floor_seeds_when_unset() {
390        let repo = setup_test_repo().await;
391        let relayer_id = uuid::Uuid::new_v4().to_string();
392        let address = uuid::Uuid::new_v4().to_string();
393
394        let effective = repo.sync_floor(&relayer_id, &address, 7).await.unwrap();
395        assert_eq!(effective, 7);
396        assert_eq!(repo.get(&relayer_id, &address).await.unwrap(), Some(7));
397    }
398
399    #[tokio::test]
400    #[ignore = "Requires active Redis instance"]
401    async fn test_get_and_increment() {
402        let repo = setup_test_repo().await;
403        let relayer_id = uuid::Uuid::new_v4().to_string();
404        let address = uuid::Uuid::new_v4().to_string();
405
406        // First increment should return 0 and set to 1
407        let result = repo.get_and_increment(&relayer_id, &address).await.unwrap();
408        assert_eq!(result, 0);
409
410        let current = repo.get(&relayer_id, &address).await.unwrap();
411        assert_eq!(current, Some(1));
412
413        // Second increment should return 1 and set to 2
414        let result = repo.get_and_increment(&relayer_id, &address).await.unwrap();
415        assert_eq!(result, 1);
416
417        let current = repo.get(&relayer_id, &address).await.unwrap();
418        assert_eq!(current, Some(2));
419    }
420
421    #[tokio::test]
422    #[ignore = "Requires active Redis instance"]
423    async fn test_decrement() {
424        let repo = setup_test_repo().await;
425        let relayer_id = uuid::Uuid::new_v4().to_string();
426        let address = uuid::Uuid::new_v4().to_string();
427
428        // Set initial value
429        repo.set(&relayer_id, &address, 5).await.unwrap();
430
431        // Decrement should return 4
432        let result = repo.decrement(&relayer_id, &address).await.unwrap();
433        assert_eq!(result, 4);
434
435        let current = repo.get(&relayer_id, &address).await.unwrap();
436        assert_eq!(current, Some(4));
437    }
438
439    #[tokio::test]
440    #[ignore = "Requires active Redis instance"]
441    async fn test_decrement_not_found() {
442        let repo = setup_test_repo().await;
443        let result = repo.decrement("nonexistent", "0x1234").await;
444        assert!(matches!(result, Err(RepositoryError::NotFound(_))));
445    }
446
447    #[tokio::test]
448    #[ignore = "Requires active Redis instance"]
449    async fn test_empty_validation() {
450        let repo = setup_test_repo().await;
451
452        // Test empty relayer_id
453        let result = repo.get("", "0x1234").await;
454        assert!(matches!(result, Err(RepositoryError::InvalidData(_))));
455
456        // Test empty address
457        let result = repo.get("relayer", "").await;
458        assert!(matches!(result, Err(RepositoryError::InvalidData(_))));
459    }
460
461    #[tokio::test]
462    #[ignore = "Requires active Redis instance"]
463    async fn test_multiple_relayers() {
464        let repo = setup_test_repo().await;
465        let relayer_1 = uuid::Uuid::new_v4().to_string();
466        let relayer_2 = uuid::Uuid::new_v4().to_string();
467        let address_1 = uuid::Uuid::new_v4().to_string();
468        let address_2 = uuid::Uuid::new_v4().to_string();
469
470        // Set different values for different relayer/address combinations
471        repo.set(&relayer_1, &address_1, 100).await.unwrap();
472        repo.set(&relayer_1, &address_2, 200).await.unwrap();
473        repo.set(&relayer_2, &address_1, 300).await.unwrap();
474
475        // Verify independent counters
476        assert_eq!(repo.get(&relayer_1, &address_1).await.unwrap(), Some(100));
477        assert_eq!(repo.get(&relayer_1, &address_2).await.unwrap(), Some(200));
478        assert_eq!(repo.get(&relayer_2, &address_1).await.unwrap(), Some(300));
479
480        // Verify independent increments
481        assert_eq!(
482            repo.get_and_increment(&relayer_1, &address_1)
483                .await
484                .unwrap(),
485            100
486        );
487        assert_eq!(
488            repo.get_and_increment(&relayer_1, &address_1)
489                .await
490                .unwrap(),
491            101
492        );
493        assert_eq!(
494            repo.get_and_increment(&relayer_1, &address_2)
495                .await
496                .unwrap(),
497            200
498        );
499        assert_eq!(
500            repo.get_and_increment(&relayer_1, &address_2)
501                .await
502                .unwrap(),
503            201
504        );
505        assert_eq!(repo.get(&relayer_2, &address_1).await.unwrap(), Some(300));
506    }
507
508    #[tokio::test]
509    #[ignore = "Requires active Redis instance"]
510    async fn test_drop_all_entries() {
511        let prefix = format!("test_drop_{}", uuid::Uuid::new_v4());
512        let repo = setup_test_repo_with_prefix(&prefix).await;
513        let relayer_1 = uuid::Uuid::new_v4().to_string();
514        let relayer_2 = uuid::Uuid::new_v4().to_string();
515        let address_1 = uuid::Uuid::new_v4().to_string();
516        let address_2 = uuid::Uuid::new_v4().to_string();
517
518        // Set up multiple counters
519        repo.set(&relayer_1, &address_1, 100).await.unwrap();
520        repo.set(&relayer_1, &address_2, 200).await.unwrap();
521        repo.set(&relayer_2, &address_1, 300).await.unwrap();
522
523        // Verify they exist
524        assert_eq!(repo.get(&relayer_1, &address_1).await.unwrap(), Some(100));
525        assert_eq!(repo.get(&relayer_1, &address_2).await.unwrap(), Some(200));
526        assert_eq!(repo.get(&relayer_2, &address_1).await.unwrap(), Some(300));
527
528        // Drop all
529        repo.drop_all_entries().await.unwrap();
530
531        // Verify all are gone
532        assert_eq!(repo.get(&relayer_1, &address_1).await.unwrap(), None);
533        assert_eq!(repo.get(&relayer_1, &address_2).await.unwrap(), None);
534        assert_eq!(repo.get(&relayer_2, &address_1).await.unwrap(), None);
535    }
536
537    #[tokio::test]
538    #[ignore = "Requires active Redis instance"]
539    async fn test_concurrent_get_and_increment() {
540        let repo = setup_test_repo().await;
541        let relayer_id = uuid::Uuid::new_v4().to_string();
542        let address = uuid::Uuid::new_v4().to_string();
543
544        // Set initial value
545        repo.set(&relayer_id, &address, 100).await.unwrap();
546
547        // Create multiple concurrent tasks that increment the counter
548        let handles: Vec<_> = (0..10)
549            .map(|_| {
550                let repo = repo.clone();
551                let relayer_id = relayer_id.clone();
552                let address = address.clone();
553                tokio::spawn(
554                    async move { repo.get_and_increment(&relayer_id, &address).await.unwrap() },
555                )
556            })
557            .collect();
558
559        // Wait for all tasks to complete and collect results
560        let mut results = Vec::new();
561        for handle in handles {
562            results.push(handle.await.unwrap());
563        }
564
565        // Sort results to check they are sequential
566        results.sort();
567
568        // Verify we get exactly the values 100-109 (no duplicates, no gaps)
569        let expected: Vec<u64> = (100..110).collect();
570        assert_eq!(results, expected);
571
572        // Verify final value is 110
573        let final_value = repo.get(&relayer_id, &address).await.unwrap();
574        assert_eq!(final_value, Some(110));
575    }
576}