use std::time::Duration;
use async_trait::async_trait;
use redis::aio::MultiplexedConnection;
use redis::{AsyncCommands, RedisError};
use crate::traits::{Cache, CacheError};
fn map_err(e: RedisError) -> CacheError {
CacheError::Io(e.to_string())
}
#[derive(Clone)]
pub struct RedisCache {
conn: MultiplexedConnection,
prefix: String,
}
impl RedisCache {
pub fn new(conn: MultiplexedConnection) -> Self {
Self::with_prefix(conn, String::new())
}
pub fn with_prefix(conn: MultiplexedConnection, prefix: String) -> Self {
Self { conn, prefix }
}
fn key(&self, key: &str) -> String {
if self.prefix.is_empty() {
key.to_string()
} else {
format!("{}{}", self.prefix, key)
}
}
}
#[async_trait]
impl Cache for RedisCache {
async fn get(&self, key: &str) -> Result<Option<Vec<u8>>, CacheError> {
let mut c = self.conn.clone();
let k = self.key(key);
let result: Result<Option<Vec<u8>>, RedisError> = c.get(&k).await;
result.map_err(map_err)
}
async fn set(
&self,
key: &str,
value: Vec<u8>,
ttl: Option<Duration>,
) -> Result<(), CacheError> {
let mut c = self.conn.clone();
let k = self.key(key);
match ttl {
Some(ttl) => {
let ms = ttl.as_millis();
if ms == 0 {
return Err(CacheError::Key(
"ttl of 0ms not allowed — pass None to store permanently".into(),
));
}
let ms = ms as u64;
let result: Result<(), RedisError> = c.pset_ex(&k, value, ms).await;
result.map_err(map_err)
}
None => {
let result: Result<(), RedisError> = c.set(&k, value).await;
result.map_err(map_err)
}
}
}
async fn invalidate(&self, key: &str) -> Result<(), CacheError> {
let mut c = self.conn.clone();
let k = self.key(key);
let result: Result<u64, RedisError> = c.del(&k).await;
result.map(|_| ()).map_err(map_err)
}
async fn clear(&self) -> Result<(), CacheError> {
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
#[ignore = "requires a live Redis instance (REDIS_URL)"]
async fn set_get_roundtrip_live() {
let url = std::env::var("REDIS_URL").expect("set REDIS_URL");
let client = redis::Client::open(url).expect("valid redis url");
let conn = client
.get_multiplexed_tokio_connection()
.await
.expect("connect");
let cache = RedisCache::with_prefix(conn, "mytheclipse_cache_test:".to_string());
cache
.set("k", b"v".to_vec(), Some(Duration::from_secs(3600)))
.await
.unwrap();
assert_eq!(cache.get("k").await.unwrap(), Some(b"v".to_vec()));
cache.invalidate("k").await.unwrap();
assert_eq!(cache.get("k").await.unwrap(), None);
}
}