use crate::backend::RedisBackend;
use crate::core::RedisCommand;
use crate::error::{OxCacheError, OxCacheResult};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
use std::time::Duration;
use tokio::sync::Mutex;
use tokio::task::JoinHandle;
const RELEASE_SCRIPT: &str = r#"
if redis.call('GET', KEYS[1]) == ARGV[1] then
return redis.call('DEL', KEYS[1])
else
return 0
end
"#;
const EXTEND_SCRIPT: &str = r#"
if redis.call('GET', KEYS[1]) == ARGV[1] then
return redis.call('PEXPIRE', KEYS[1], ARGV[2])
else
return 0
end
"#;
pub struct DistributedLock {
pub(super) backend: Arc<RedisBackend>,
pub(super) key: String,
pub(super) owner_id: String,
pub(super) reentrant_count: AtomicU32,
pub(super) ttl: Duration,
pub(super) watchdog_enabled: bool,
pub(super) watchdog: Mutex<Option<JoinHandle<()>>>,
pub(super) released: Arc<AtomicBool>,
}
fn watchdog_retry_delay(consecutive_errors: u32) -> Duration {
let shift = consecutive_errors.min(5);
Duration::from_millis(200 * (1 << shift))
}
impl DistributedLock {
pub async fn acquire(&mut self) -> OxCacheResult<bool> {
let count = self.reentrant_count.load(Ordering::SeqCst);
if count > 0 {
self.reentrant_count.fetch_add(1, Ordering::SeqCst);
return Ok(false);
}
self.released.store(false, Ordering::SeqCst);
let ttl_ms = self.ttl.as_millis() as u64;
let mut conn = self.backend.conn();
let result: Option<String> = redis::cmd(RedisCommand::Set.as_str())
.arg(&self.key)
.arg(&self.owner_id)
.arg("NX")
.arg("PX")
.arg(ttl_ms)
.query_async(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("dist_lock acquire failed: {e}")))?;
match result {
Some(_) => {
self.reentrant_count.store(1, Ordering::SeqCst);
if self.watchdog_enabled {
let handle = self.spawn_watchdog();
*self.watchdog.lock().await = Some(handle);
}
Ok(true)
}
None => {
Err(OxCacheError::Operation(format!(
"dist_lock '{}' is already held by another owner",
self.key
)))
}
}
}
pub async fn release(&mut self) -> OxCacheResult<()> {
let count = self.reentrant_count.load(Ordering::SeqCst);
if count == 0 {
return Err(OxCacheError::Operation(
"dist_lock not held, cannot release".to_string(),
));
}
let new_count = count - 1;
if new_count > 0 {
self.reentrant_count.store(new_count, Ordering::SeqCst);
return Ok(());
}
let mut conn = self.backend.conn();
let result: i64 = redis::cmd(RedisCommand::Eval.as_str())
.arg(RELEASE_SCRIPT)
.arg(1)
.arg(&self.key)
.arg(&self.owner_id)
.query_async(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("dist_lock release failed: {e}")))?;
if result == 0 {
return Err(OxCacheError::Operation(format!(
"dist_lock '{}' not held by owner (already expired or stolen)",
self.key
)));
}
self.reentrant_count.store(0, Ordering::SeqCst);
self.released.store(true, Ordering::SeqCst);
if let Some(handle) = self.watchdog.lock().await.take() {
handle.abort();
}
Ok(())
}
pub async fn extend(&self) -> OxCacheResult<bool> {
let ttl_ms = self.ttl.as_millis().to_string();
let mut conn = self.backend.conn();
let result: i64 = redis::cmd(RedisCommand::Eval.as_str())
.arg(EXTEND_SCRIPT)
.arg(1)
.arg(&self.key)
.arg(&self.owner_id)
.arg(&ttl_ms)
.query_async(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("dist_lock extend failed: {e}")))?;
Ok(result == 1)
}
pub async fn is_held(&self) -> OxCacheResult<bool> {
let mut conn = self.backend.conn();
let result: Option<String> = redis::cmd(RedisCommand::Get.as_str())
.arg(&self.key)
.query_async(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("dist_lock is_held check failed: {e}")))?;
Ok(result.as_deref() == Some(&self.owner_id))
}
fn spawn_watchdog(&self) -> JoinHandle<()> {
let backend = self.backend.clone();
let key = self.key.clone();
let owner_id = self.owner_id.clone();
let ttl = self.ttl;
let released = self.released.clone();
let renew_interval = ttl / 3;
tokio::spawn(async move {
let mut consecutive_errors: u32 = 0;
loop {
tokio::time::sleep(renew_interval).await;
if released.load(Ordering::SeqCst) {
break;
}
let ttl_ms = ttl.as_millis().to_string();
let mut conn = backend.conn();
let result: Result<i64, _> = redis::cmd(RedisCommand::Eval.as_str())
.arg(EXTEND_SCRIPT)
.arg(1)
.arg(&key)
.arg(&owner_id)
.arg(&ttl_ms)
.query_async(&mut conn)
.await;
match result {
Ok(1) => {
consecutive_errors = 0;
}
Ok(_) => {
break;
}
Err(_) => {
consecutive_errors = consecutive_errors.saturating_add(1);
tokio::time::sleep(watchdog_retry_delay(consecutive_errors)).await;
}
}
}
})
}
}
impl Drop for DistributedLock {
fn drop(&mut self) {
self.released.store(true, Ordering::SeqCst);
}
}
#[cfg(test)]
mod tests {
use super::*;
use uuid::Uuid;
#[test]
fn test_reentrant_count_logic() {
let count = AtomicU32::new(0);
assert_eq!(count.load(Ordering::SeqCst), 0);
count.fetch_add(1, Ordering::SeqCst);
assert_eq!(count.load(Ordering::SeqCst), 1);
count.fetch_add(1, Ordering::SeqCst);
assert_eq!(count.load(Ordering::SeqCst), 2);
let new_count = count.fetch_sub(1, Ordering::SeqCst) - 1;
assert_eq!(new_count, 1);
assert!(new_count > 0);
let new_count = count.fetch_sub(1, Ordering::SeqCst) - 1;
assert_eq!(new_count, 0);
}
#[test]
fn test_owner_id_is_unique() {
let id1 = Uuid::new_v4().to_string();
let id2 = Uuid::new_v4().to_string();
assert_ne!(id1, id2);
}
#[test]
fn test_lua_scripts_are_valid() {
assert!(RELEASE_SCRIPT.contains("GET"));
assert!(RELEASE_SCRIPT.contains("DEL"));
assert!(EXTEND_SCRIPT.contains("GET"));
assert!(EXTEND_SCRIPT.contains("PEXPIRE"));
}
#[test]
fn test_released_flag_default() {
let released = AtomicBool::new(false);
assert!(!released.load(Ordering::SeqCst));
released.store(true, Ordering::SeqCst);
assert!(released.load(Ordering::SeqCst));
}
#[test]
fn test_watchdog_retry_delay_grows_and_caps() {
assert!(watchdog_retry_delay(0) <= watchdog_retry_delay(1));
assert!(watchdog_retry_delay(1) <= watchdog_retry_delay(5));
assert!(watchdog_retry_delay(u32::MAX) <= Duration::from_secs(10));
}
#[allow(unsafe_code)]
async fn live_test_backend() -> Arc<RedisBackend> {
unsafe {
std::env::set_var("OXCACHE_ALLOW_INSECURE_REDIS", "I_UNDERSTAND_THE_RISKS");
}
Arc::new(
RedisBackend::new("redis://127.0.0.1:6379")
.await
.expect("live Redis required at 127.0.0.1:6379"),
)
}
#[tokio::test]
#[ignore = "needs live Redis at 127.0.0.1:6379"]
async fn test_release_stolen_lock_keeps_retryable_state() {
use super::super::DistLockBuilder;
let backend = live_test_backend().await;
let key = format!("test-lock-stolen-{}", Uuid::new_v4());
let mut lock = DistLockBuilder::new(backend.clone(), key.clone())
.ttl(Duration::from_secs(30))
.watchdog_enabled(false)
.build();
assert!(lock.acquire().await.expect("acquire"));
let mut conn = backend.conn();
let _: i64 = redis::cmd(RedisCommand::Del.as_str())
.arg(&key)
.query_async(&mut conn)
.await
.expect("DEL");
let err = lock
.release()
.await
.expect_err("release of stolen lock must fail");
assert!(err.to_string().contains("not held by owner"));
assert_eq!(
lock.reentrant_count.load(Ordering::SeqCst),
1,
"failed release must keep retryable state"
);
}
#[tokio::test]
#[ignore = "needs live Redis at 127.0.0.1:6379"]
async fn test_release_conn_error_keeps_retryable_state() {
use super::super::DistLockBuilder;
use std::process::Command;
let backend = live_test_backend().await;
let key = format!("test-lock-connerr-{}", Uuid::new_v4());
let mut lock = DistLockBuilder::new(backend.clone(), key.clone())
.ttl(Duration::from_secs(30))
.watchdog_enabled(false)
.build();
assert!(lock.acquire().await.expect("acquire"));
let _ = Command::new("redis-cli")
.args(["-p", "6379", "shutdown", "nosave"])
.status();
tokio::time::sleep(Duration::from_millis(300)).await;
let err = lock
.release()
.await
.expect_err("release against dead server must fail");
assert!(err.to_string().contains("dist_lock release failed"));
assert_eq!(
lock.reentrant_count.load(Ordering::SeqCst),
1,
"failed release must keep retryable state"
);
assert!(
!lock.released.load(Ordering::SeqCst),
"failed release must not flag the lock as released"
);
let _ = Command::new("redis-server")
.args([
"--port",
"6379",
"--daemonize",
"yes",
"--save",
"",
"--appendonly",
"no",
])
.status();
tokio::time::sleep(Duration::from_millis(500)).await;
}
#[tokio::test]
#[ignore = "needs live Redis at 127.0.0.1:6379"]
async fn test_watchdog_renews_past_ttl() {
use super::super::DistLockBuilder;
let backend = live_test_backend().await;
let key = format!("test-lock-watchdog-{}", Uuid::new_v4());
let mut lock = DistLockBuilder::new(backend, key)
.ttl(Duration::from_secs(3))
.watchdog_enabled(true)
.build();
assert!(lock.acquire().await.expect("acquire"));
tokio::time::sleep(Duration::from_millis(4500)).await;
assert!(
lock.is_held().await.expect("is_held"),
"watchdog must renew the lock past TTL"
);
lock.release().await.expect("release");
}
}