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>,
}
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;
self.reentrant_count.store(new_count, Ordering::SeqCst);
if new_count > 0 {
return Ok(());
}
self.released.store(true, Ordering::SeqCst);
if let Some(handle) = self.watchdog.lock().await.take() {
handle.abort();
}
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
)));
}
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 {
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) => {
}
Ok(_) => {
break;
}
Err(_) => {
break;
}
}
}
})
}
}
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));
}
}