use crate::error::Result;
use fs2::FileExt;
use std::fs::{self, File, OpenOptions};
use std::path::{Path, PathBuf};
#[derive(Debug, Clone)]
pub struct RefreshLockManager {
lock_dir: PathBuf,
}
impl RefreshLockManager {
pub fn new(lock_dir: PathBuf) -> Result<Self> {
fs::create_dir_all(&lock_dir)?;
Ok(Self { lock_dir })
}
pub fn with_default_dir() -> Result<Self> {
let lock_dir = Self::default_lock_dir()?;
Self::new(lock_dir)
}
pub fn for_app(app_name: &str) -> Result<Self> {
let mut lock_dir = Self::default_lock_dir()?;
lock_dir.push(app_name);
Self::new(lock_dir)
}
fn default_lock_dir() -> Result<PathBuf> {
if let Ok(runtime_dir) = std::env::var("XDG_RUNTIME_DIR") {
let mut path = PathBuf::from(runtime_dir);
path.push("schlussel-locks");
return Ok(path);
}
let mut path = std::env::temp_dir();
path.push(format!("schlussel-locks-{}", Self::get_user_id()));
Ok(path)
}
#[cfg(unix)]
fn get_user_id() -> String {
use std::os::unix::fs::MetadataExt;
std::env::current_exe()
.and_then(|p| p.metadata())
.map(|m| m.uid().to_string())
.unwrap_or_else(|_| "unknown".to_string())
}
#[cfg(not(unix))]
fn get_user_id() -> String {
std::env::var("USERNAME")
.or_else(|_| std::env::var("USER"))
.unwrap_or_else(|_| "unknown".to_string())
}
pub fn acquire_lock(&self, key: &str) -> Result<RefreshLock> {
let lock_path = self.lock_path(key);
if let Some(parent) = lock_path.parent() {
fs::create_dir_all(parent)?;
}
let file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(false)
.open(&lock_path)?;
file.lock_exclusive()?;
Ok(RefreshLock {
file: Some(file),
path: lock_path,
})
}
pub fn try_acquire_lock(&self, key: &str) -> Result<Option<RefreshLock>> {
let lock_path = self.lock_path(key);
if let Some(parent) = lock_path.parent() {
fs::create_dir_all(parent)?;
}
let file = OpenOptions::new()
.read(true)
.write(true)
.create(true)
.truncate(false)
.open(&lock_path)?;
match file.try_lock_exclusive() {
Ok(()) => Ok(Some(RefreshLock {
file: Some(file),
path: lock_path,
})),
Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => Ok(None),
#[cfg(windows)]
Err(e) if e.raw_os_error() == Some(33) => Ok(None),
Err(e) => Err(e.into()),
}
}
fn lock_path(&self, key: &str) -> PathBuf {
let safe_key = key.replace(['/', '\\', ':', '*', '?', '"', '<', '>', '|'], "_");
self.lock_dir.join(format!("{}.lock", safe_key))
}
}
pub struct RefreshLock {
file: Option<File>,
path: PathBuf,
}
impl RefreshLock {
pub fn path(&self) -> &Path {
&self.path
}
}
impl Drop for RefreshLock {
fn drop(&mut self) {
if let Some(file) = self.file.take() {
let _ = file.unlock();
}
let _ = fs::remove_file(&self.path);
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
use std::thread;
use std::time::Duration;
#[test]
fn test_lock_manager_creation() {
let temp_dir = std::env::temp_dir().join(format!("test_locks_{}", rand::random::<u32>()));
let _manager = RefreshLockManager::new(temp_dir.clone()).unwrap();
assert!(temp_dir.exists());
fs::remove_dir_all(temp_dir).ok();
}
#[test]
fn test_acquire_and_release_lock() {
let temp_dir = std::env::temp_dir().join(format!("test_locks_{}", rand::random::<u32>()));
let manager = RefreshLockManager::new(temp_dir.clone()).unwrap();
let lock = manager.acquire_lock("test-key").unwrap();
assert!(lock.path().exists());
drop(lock);
let lock2 = manager.acquire_lock("test-key").unwrap();
drop(lock2);
fs::remove_dir_all(temp_dir).ok();
}
#[test]
fn test_concurrent_lock_attempts() {
let temp_dir = std::env::temp_dir().join(format!("test_locks_{}", rand::random::<u32>()));
let manager = Arc::new(RefreshLockManager::new(temp_dir.clone()).unwrap());
let manager1 = manager.clone();
let manager2 = manager.clone();
let handle1 = thread::spawn(move || {
let _lock = manager1.acquire_lock("concurrent-test").unwrap();
thread::sleep(Duration::from_millis(200));
"thread1"
});
thread::sleep(Duration::from_millis(50));
let handle2 = thread::spawn(move || {
let _lock = manager2.acquire_lock("concurrent-test").unwrap();
"thread2"
});
let result1 = handle1.join().unwrap();
let result2 = handle2.join().unwrap();
assert_eq!(result1, "thread1");
assert_eq!(result2, "thread2");
fs::remove_dir_all(temp_dir).ok();
}
#[test]
fn test_try_acquire_lock() {
let temp_dir = std::env::temp_dir().join(format!("test_locks_{}", rand::random::<u32>()));
let manager = RefreshLockManager::new(temp_dir.clone()).unwrap();
let lock1 = manager.try_acquire_lock("try-test").unwrap();
assert!(lock1.is_some());
let lock2 = manager.try_acquire_lock("try-test").unwrap();
assert!(lock2.is_none());
drop(lock1);
let lock3 = manager.try_acquire_lock("try-test").unwrap();
assert!(lock3.is_some());
fs::remove_dir_all(temp_dir).ok();
}
#[test]
fn test_key_sanitization() {
let temp_dir = std::env::temp_dir().join(format!("test_locks_{}", rand::random::<u32>()));
let manager = RefreshLockManager::new(temp_dir.clone()).unwrap();
let lock = manager.acquire_lock("domain.com:user/name").unwrap();
assert!(lock
.path()
.to_str()
.unwrap()
.contains("domain.com_user_name.lock"));
fs::remove_dir_all(temp_dir).ok();
}
}