1use async_trait::async_trait;
2use redis::{Client, AsyncCommands};
3use std::time::Duration;
4use serde::{Deserialize, Serialize};
5use crate::{validate_cache_key, validate_ttl, Cache, Result};
6
7pub struct RedisCache {
9 client: Client,
10 default_ttl: Option<Duration>,
11}
12
13impl RedisCache {
14 pub fn new(url: &str) -> Result<Self> {
16 let client = Client::open(url)?;
17
18 Ok(Self {
19 client,
20 default_ttl: Some(Duration::from_secs(3600)),
21 })
22 }
23
24 pub fn with_default_ttl(url: &str, ttl: Duration) -> Result<Self> {
26 validate_ttl(Some(ttl))?;
27 let client = Client::open(url)?;
28
29 Ok(Self {
30 client,
31 default_ttl: Some(ttl),
32 })
33 }
34}
35
36#[async_trait]
37impl Cache for RedisCache {
38 async fn get<T>(&self, key: &str) -> Result<Option<T>>
39 where
40 T: for<'de> Deserialize<'de> + Send,
41 {
42 validate_cache_key(key)?;
43 let mut conn = self.client.get_multiplexed_async_connection().await?;
44
45 let result: Option<String> = conn.get(key).await?;
46
47 if let Some(data) = result {
48 let value: T = serde_json::from_str(&data)?;
49 Ok(Some(value))
50 } else {
51 Ok(None)
52 }
53 }
54
55 async fn set<T>(&self, key: &str, value: &T, ttl: Option<Duration>) -> Result<()>
56 where
57 T: Serialize + Send + Sync,
58 {
59 validate_cache_key(key)?;
60 validate_ttl(ttl)?;
61 let mut conn = self.client.get_multiplexed_async_connection().await?;
62
63 let data = serde_json::to_string(value)?;
64
65 let ttl = ttl.or(self.default_ttl);
66
67 if let Some(duration) = ttl {
68 let seconds = duration.as_secs().max(1);
69 let _: () = conn.set_ex(key, data, seconds).await?;
70 } else {
71 let _: () = conn.set(key, data).await?;
72 }
73
74 Ok(())
75 }
76
77 async fn delete(&self, key: &str) -> Result<()> {
78 validate_cache_key(key)?;
79 let mut conn = self.client.get_multiplexed_async_connection().await?;
80
81 let _: () = conn.del(key).await?;
82
83 Ok(())
84 }
85
86 async fn exists(&self, key: &str) -> Result<bool> {
87 validate_cache_key(key)?;
88 let mut conn = self.client.get_multiplexed_async_connection().await?;
89
90 let exists: bool = conn.exists(key).await?;
91
92 Ok(exists)
93 }
94
95 async fn flush(&self) -> Result<()> {
96 let mut conn = self.client.get_multiplexed_async_connection().await?;
97
98 let _: () = redis::cmd("FLUSHDB")
99 .query_async(&mut conn)
100 .await?;
101
102 Ok(())
103 }
104}
105
106#[cfg(test)]
107mod tests {
108 use super::RedisCache;
109 use std::time::Duration;
110
111 #[test]
112 fn rejects_zero_default_ttl() {
113 let result = RedisCache::with_default_ttl("redis://127.0.0.1/", Duration::from_secs(0));
114 assert!(result.is_err());
115 }
116}