use crate::backend::RedisBackend;
use crate::core::RedisCommand;
use crate::error::{OxCacheError, OxCacheResult};
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicU32, AtomicU64, 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>,
pub(super) fencing_token: AtomicU64,
}
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);
match self.acquire_fence_token().await {
Ok(token) => {
self.fencing_token.store(token, Ordering::SeqCst);
}
Err(_) => {
}
}
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
)))
}
}
}
async fn acquire_fence_token(&self) -> OxCacheResult<u64> {
let mut conn = self.backend.conn();
let token: i64 = redis::cmd(RedisCommand::Incr.as_str())
.arg(format!("{}:fence", self.key))
.query_async(&mut conn)
.await
.map_err(|e| OxCacheError::Operation(format!("dist_lock fence incr failed: {e}")))?;
Ok(u64::try_from(token).unwrap_or(0))
}
pub fn token(&self) -> u64 {
self.fencing_token.load(Ordering::SeqCst)
}
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;