1use 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 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 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 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 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 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 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 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 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 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 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 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 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 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 repo.set(&relayer_id, &address, 5).await.unwrap();
430
431 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 let result = repo.get("", "0x1234").await;
454 assert!(matches!(result, Err(RepositoryError::InvalidData(_))));
455
456 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 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 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 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 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 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 repo.drop_all_entries().await.unwrap();
530
531 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 repo.set(&relayer_id, &address, 100).await.unwrap();
546
547 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 let mut results = Vec::new();
561 for handle in handles {
562 results.push(handle.await.unwrap());
563 }
564
565 results.sort();
567
568 let expected: Vec<u64> = (100..110).collect();
570 assert_eq!(results, expected);
571
572 let final_value = repo.get(&relayer_id, &address).await.unwrap();
574 assert_eq!(final_value, Some(110));
575 }
576}