use std::collections::HashMap;
use std::time::Duration;
use async_trait::async_trait;
use fred::interfaces::LuaInterface;
use fred::prelude::*;
use fred::types::ConnectHandle;
use crate::backend::{Backend, HealthStatus, LockableBackend, TtlInspectable};
use crate::error::{BackendError, BackendErrorKind};
fn sanitize_redis_message(msg: &str) -> String {
if let Some(proto_end) = msg.find("://") {
if let Some(at_pos) = msg[proto_end..].find('@') {
let mut sanitized = String::with_capacity(msg.len());
sanitized.push_str(&msg[..proto_end + 3]);
sanitized.push_str("[REDACTED]");
sanitized.push_str(&msg[proto_end + at_pos..]);
return sanitized;
}
}
msg.to_string()
}
fn redis_err(e: RedisError) -> BackendError {
let kind = match e.kind() {
RedisErrorKind::Auth => BackendErrorKind::Authentication,
RedisErrorKind::IO => BackendErrorKind::Transient,
RedisErrorKind::Timeout => BackendErrorKind::Timeout,
RedisErrorKind::Canceled => BackendErrorKind::Transient,
_ => BackendErrorKind::Permanent,
};
BackendError {
kind,
message: sanitize_redis_message(&e.to_string()),
source: Some(Box::new(e)),
}
}
pub struct RedisBackend {
client: RedisClient,
}
impl std::fmt::Debug for RedisBackend {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("RedisBackend").finish_non_exhaustive()
}
}
impl RedisBackend {
pub fn builder() -> RedisBackendBuilder {
RedisBackendBuilder::default()
}
pub async fn connect(&self) -> Result<ConnectHandle, BackendError> {
self.client.init().await.map_err(redis_err)
}
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(not(feature = "unsync"), async_trait)]
#[cfg_attr(feature = "unsync", async_trait(?Send))]
impl Backend for RedisBackend {
async fn get(&self, key: &str) -> Result<Option<Vec<u8>>, BackendError> {
let result: Option<bytes::Bytes> = self.client.get(key).await.map_err(redis_err)?;
Ok(result.map(|b| b.to_vec()))
}
async fn set(
&self,
key: &str,
value: Vec<u8>,
ttl: Option<Duration>,
) -> Result<(), BackendError> {
let expiration = ttl.map(|d| {
let secs = i64::try_from(d.as_secs().max(1)).unwrap_or(i64::MAX);
Expiration::EX(secs)
});
self.client
.set::<(), _, _>(key, value.as_slice(), expiration, None, false)
.await
.map_err(redis_err)
}
async fn delete(&self, key: &str) -> Result<bool, BackendError> {
let removed: i64 = self.client.del(key).await.map_err(redis_err)?;
Ok(removed > 0)
}
async fn exists(&self, key: &str) -> Result<bool, BackendError> {
let count: i64 = self.client.exists(key).await.map_err(redis_err)?;
Ok(count > 0)
}
async fn health(&self) -> Result<HealthStatus, BackendError> {
let start = std::time::Instant::now();
let _pong: String = self.client.ping().await.map_err(redis_err)?;
let latency = start.elapsed();
let mut details = HashMap::new();
details.insert("latency_ms".to_string(), latency.as_millis().to_string());
Ok(HealthStatus {
is_healthy: true,
latency_ms: latency.as_secs_f64() * 1000.0,
backend_type: "redis".to_string(),
details,
})
}
}
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(not(feature = "unsync"), async_trait)]
#[cfg_attr(feature = "unsync", async_trait(?Send))]
impl TtlInspectable for RedisBackend {
async fn ttl(&self, key: &str) -> Result<Option<Duration>, BackendError> {
let secs: i64 = self.client.ttl(key).await.map_err(redis_err)?;
match secs {
..0 => Ok(None),
n => Ok(Some(Duration::from_secs(n.unsigned_abs()))),
}
}
}
fn lock_key(key: &str) -> String {
format!("{key}:lock")
}
const RELEASE_LOCK_SCRIPT: &str = r#"if redis.call("GET", KEYS[1]) == ARGV[1] then return redis.call("DEL", KEYS[1]) else return 0 end"#;
#[cfg(not(target_arch = "wasm32"))]
#[cfg_attr(not(feature = "unsync"), async_trait)]
#[cfg_attr(feature = "unsync", async_trait(?Send))]
impl LockableBackend for RedisBackend {
async fn acquire_lock(
&self,
key: &str,
timeout_ms: u64,
) -> Result<Option<String>, BackendError> {
let lock_id = uuid::Uuid::new_v4().to_string();
let px = i64::try_from(timeout_ms.max(1)).unwrap_or(i64::MAX);
let result: Option<String> = self
.client
.set(
lock_key(key),
lock_id.as_str(),
Some(Expiration::PX(px)),
Some(SetOptions::NX),
false,
)
.await
.map_err(redis_err)?;
Ok(result.map(|_| lock_id))
}
async fn release_lock(&self, key: &str, lock_id: &str) -> Result<bool, BackendError> {
let deleted: i64 = self
.client
.eval(RELEASE_LOCK_SCRIPT, vec![lock_key(key)], vec![lock_id])
.await
.map_err(redis_err)?;
Ok(deleted == 1)
}
}
#[derive(Default)]
#[must_use]
pub struct RedisBackendBuilder {
url: Option<String>,
reconnect: Option<ReconnectPolicy>,
}
impl RedisBackendBuilder {
pub fn url(mut self, url: impl Into<String>) -> Self {
self.url = Some(url.into());
self
}
pub(crate) fn auto_reconnect(mut self) -> Self {
self.reconnect = Some(ReconnectPolicy::new_exponential(0, 100, 30_000, 2));
self
}
pub fn build(self) -> Result<RedisBackend, crate::error::CachekitError> {
use crate::error::CachekitError;
let url = self
.url
.filter(|u| !u.is_empty())
.ok_or_else(|| CachekitError::Config("url is required".to_string()))?;
let config = RedisConfig::from_url(&url).map_err(|e| {
CachekitError::Config(format!(
"invalid Redis URL: {}",
sanitize_redis_message(&e.to_string())
))
})?;
let client = RedisClient::new(config, None, None, self.reconnect);
Ok(RedisBackend { client })
}
}
#[cfg(test)]
mod tests {
use super::*;
fn _assert_lockable(_b: &dyn LockableBackend) {}
#[test]
fn redis_is_lockable() {
fn _check(backend: &RedisBackend) {
_assert_lockable(backend);
}
}
#[test]
fn lock_key_derives_py_compatible_namespace() {
assert_eq!(
lock_key("ns:app:func:m.f:args:abc:v1"),
"ns:app:func:m.f:args:abc:v1:lock"
);
}
}