use std::future::Future;
use std::time::Duration;
use redis::AsyncCommands;
use crate::cache::config::CacheConfig;
use crate::cache::error::{CacheConnectError, CacheError, CacheHealthError};
use crate::cache::namespace::Namespace;
#[derive(Clone)]
pub struct Cache {
pub(crate) connection: redis::aio::MultiplexedConnection,
pub(crate) response_timeout: Duration,
pub(crate) max_payload_size: Option<usize>,
pub(crate) namespace: Namespace,
}
impl Cache {
pub(crate) fn from_parts(
connection: redis::aio::MultiplexedConnection,
response_timeout: Duration,
max_payload_size: Option<usize>,
namespace: Namespace,
) -> Self {
Self {
connection,
response_timeout,
max_payload_size,
namespace,
}
}
pub(crate) fn connection_for_op(&self) -> redis::aio::MultiplexedConnection {
let mut conn = self.connection.clone();
conn.set_response_timeout(self.response_timeout);
conn
}
pub(crate) fn resolve_key(&self, key: &str) -> String {
self.namespace.resolve(key)
}
#[must_use]
pub fn connection(&self) -> redis::aio::MultiplexedConnection {
self.connection.clone()
}
}
impl Cache {
pub async fn connect(config: CacheConfig) -> Result<Cache, CacheConnectError> {
config.validate().map_err(CacheConnectError::config)?;
let connection = config
.client()
.get_multiplexed_async_connection()
.await
.map_err(CacheConnectError::backend)?;
Ok(Cache::from_parts(
connection,
config.response_timeout_setting(),
config.max_payload_size_setting(),
config.namespace_setting().clone(),
))
}
#[must_use]
pub fn from_connection(
config: CacheConfig,
connection: redis::aio::MultiplexedConnection,
) -> Cache {
Cache::from_parts(
connection,
config.response_timeout_setting(),
config.max_payload_size_setting(),
config.namespace_setting().clone(),
)
}
pub async fn ping(&self) -> Result<(), CacheHealthError> {
let mut conn = self.connection_for_op();
redis::cmd("PING")
.query_async::<String>(&mut conn)
.await
.map(|_| ())
.map_err(CacheHealthError::new)
}
pub async fn close(&self) {
}
}
impl Cache {
pub async fn get_bytes(&self, key: &str) -> Result<Option<Vec<u8>>, CacheError> {
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
let value: Option<Vec<u8>> = conn.get(&full_key).await.map_err(CacheError::backend)?;
if let Some(ref bytes) = value {
check_payload_size(bytes.len(), self.max_payload_size)?;
}
Ok(value)
}
pub async fn get<T: serde::de::DeserializeOwned>(
&self,
key: &str,
) -> Result<Option<T>, CacheError> {
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
let value: Option<Vec<u8>> = conn.get(&full_key).await.map_err(CacheError::backend)?;
match value {
None => Ok(None),
Some(bytes) => {
check_payload_size(bytes.len(), self.max_payload_size)?;
let parsed = serde_json::from_slice(&bytes)
.map_err(|source| CacheError::Decode { source })?;
Ok(Some(parsed))
}
}
}
pub async fn set_bytes(&self, key: &str, value: Vec<u8>) -> Result<(), CacheError> {
check_payload_size(value.len(), self.max_payload_size)?;
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
conn.set::<_, _, ()>(full_key, value)
.await
.map_err(CacheError::backend)
}
pub async fn set<T: serde::Serialize>(&self, key: &str, value: &T) -> Result<(), CacheError> {
let bytes = serde_json::to_vec(value).map_err(|source| CacheError::Decode { source })?;
self.set_bytes(key, bytes).await
}
pub async fn put<T: serde::Serialize>(
&self,
key: &str,
value: &T,
ttl: Duration,
) -> Result<(), CacheError> {
if ttl.is_zero() {
return Err(CacheError::ZeroTtl);
}
let bytes = serde_json::to_vec(value).map_err(|source| CacheError::Decode { source })?;
check_payload_size(bytes.len(), self.max_payload_size)?;
let full_key = self.resolve_key(key);
let seconds = ttl.as_secs();
let mut conn = self.connection_for_op();
conn.set_ex::<_, _, ()>(full_key, bytes, seconds)
.await
.map_err(CacheError::backend)
}
pub async fn set_bytes_with_ttl(
&self,
key: &str,
value: Vec<u8>,
ttl: Duration,
) -> Result<(), CacheError> {
if ttl.is_zero() {
return Err(CacheError::ZeroTtl);
}
check_payload_size(value.len(), self.max_payload_size)?;
let full_key = self.resolve_key(key);
let seconds = ttl.as_secs();
let mut conn = self.connection_for_op();
conn.set_ex::<_, _, ()>(full_key, value, seconds)
.await
.map_err(CacheError::backend)
}
pub async fn remember<T, F, Fut, E>(
&self,
key: &str,
ttl: Duration,
loader: F,
) -> Result<T, CacheError>
where
T: serde::Serialize + serde::de::DeserializeOwned,
F: FnOnce() -> Fut,
Fut: Future<Output = Result<T, E>>,
E: Into<CacheError>,
{
match self.get::<T>(key).await? {
Some(value) => Ok(value),
None => {
let value = loader().await.map_err(E::into)?;
self.put(key, &value, ttl).await?;
Ok(value)
}
}
}
pub async fn forget(&self, key: &str) -> Result<usize, CacheError> {
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
conn.del::<_, usize>(full_key)
.await
.map_err(CacheError::backend)
}
pub async fn exists(&self, key: &str) -> Result<bool, CacheError> {
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
let count: usize = conn
.exists::<_, usize>(full_key)
.await
.map_err(CacheError::backend)?;
Ok(count > 0)
}
pub async fn incr(&self, key: &str, delta: i64) -> Result<i64, CacheError> {
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
conn.incr::<_, _, i64>(full_key, delta)
.await
.map_err(CacheError::backend)
}
pub async fn decr(&self, key: &str, delta: i64) -> Result<i64, CacheError> {
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
conn.decr::<_, _, i64>(full_key, delta)
.await
.map_err(CacheError::backend)
}
pub async fn expire(&self, key: &str, ttl: Duration) -> Result<bool, CacheError> {
if ttl.is_zero() {
return Err(CacheError::ZeroTtl);
}
let full_key = self.resolve_key(key);
let seconds: i64 = ttl.as_secs() as i64;
let mut conn = self.connection_for_op();
let set: bool = conn
.expire::<_, bool>(full_key, seconds)
.await
.map_err(CacheError::backend)?;
Ok(set)
}
pub async fn ttl(&self, key: &str) -> Result<Option<u64>, CacheError> {
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
let raw: i64 = conn
.ttl::<_, i64>(full_key)
.await
.map_err(CacheError::backend)?;
match raw {
-1 => Ok(None),
-2 => Ok(Some(0)),
secs if secs >= 0 => Ok(Some(secs as u64)),
_ => Ok(Some(0)),
}
}
pub async fn set_if_absent(&self, key: &str, value: Vec<u8>) -> Result<bool, CacheError> {
check_payload_size(value.len(), self.max_payload_size)?;
let full_key = self.resolve_key(key);
let mut conn = self.connection_for_op();
let result: Option<String> = redis::cmd("SET")
.arg(&full_key)
.arg(value)
.arg("NX")
.query_async(&mut conn)
.await
.map_err(CacheError::backend)?;
Ok(result.is_some())
}
pub async fn set_if_absent_with_ttl(
&self,
key: &str,
value: Vec<u8>,
ttl: Duration,
) -> Result<bool, CacheError> {
if ttl.is_zero() {
return Err(CacheError::ZeroTtl);
}
check_payload_size(value.len(), self.max_payload_size)?;
let full_key = self.resolve_key(key);
let seconds = ttl.as_secs();
let mut conn = self.connection_for_op();
let result: Option<String> = redis::cmd("SET")
.arg(&full_key)
.arg(value)
.arg("NX")
.arg("EX")
.arg(seconds)
.query_async(&mut conn)
.await
.map_err(CacheError::backend)?;
Ok(result.is_some())
}
}
pub(crate) fn check_payload_size(size: usize, limit: Option<usize>) -> Result<(), CacheError> {
if let Some(limit) = limit
&& size > limit
{
return Err(CacheError::PayloadTooLarge { size, limit });
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn no_limit_accepts_anything() {
assert!(check_payload_size(1_000_000_000, None).is_ok());
}
#[test]
fn within_limit_passes() {
assert!(check_payload_size(100, Some(1024)).is_ok());
}
#[test]
fn exactly_at_limit_passes() {
assert!(check_payload_size(1024, Some(1024)).is_ok());
}
#[test]
fn over_limit_fails() {
let result = check_payload_size(1025, Some(1024));
assert!(matches!(
result,
Err(CacheError::PayloadTooLarge {
size: 1025,
limit: 1024
})
));
}
#[test]
fn zero_size_passes() {
assert!(check_payload_size(0, Some(1024)).is_ok());
}
#[tokio::test]
#[ignore = "needs a live Redis on 127.0.0.1:6379"]
async fn a_failing_loader_does_not_poison_the_cache() {
use crate::cache::Namespace;
use std::sync::atomic::{AtomicUsize, Ordering};
let namespace =
Namespace::new(&format!("remember-test-{}", std::process::id())).expect("a namespace");
let cache = Cache::connect(
CacheConfig::new("redis://127.0.0.1:6379")
.expect("a cache config")
.namespace(namespace),
)
.await
.expect("a live Redis");
let ttl = Duration::from_secs(60);
let calls = AtomicUsize::new(0);
let failed: Result<String, CacheError> = cache
.remember("greeting", ttl, || {
calls.fetch_add(1, Ordering::Relaxed);
async { Err(CacheError::ZeroTtl) }
})
.await;
assert!(failed.is_err());
assert_eq!(calls.load(Ordering::Relaxed), 1);
assert_eq!(
cache.get::<String>("greeting").await.expect("a get"),
None,
"a failed load must not leave anything behind"
);
let loaded: String = cache
.remember("greeting", ttl, || {
calls.fetch_add(1, Ordering::Relaxed);
async { Ok::<_, CacheError>("hello".to_string()) }
})
.await
.expect("a successful load");
assert_eq!(loaded, "hello");
assert_eq!(calls.load(Ordering::Relaxed), 2);
let cached: String = cache
.remember("greeting", ttl, || {
calls.fetch_add(1, Ordering::Relaxed);
async { Err(CacheError::ZeroTtl) }
})
.await
.expect("a cache hit");
assert_eq!(cached, "hello");
assert_eq!(
calls.load(Ordering::Relaxed),
2,
"the loader must not run on a hit"
);
cache.forget("greeting").await.expect("cleanup");
}
}